mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
Merge remote-tracking branch 'origin/main' into litellm_bulk_user_delete
This commit is contained in:
commit
ef71222349
135 changed files with 7013 additions and 1551 deletions
|
|
@ -5,6 +5,7 @@ use serde_json::Value;
|
|||
|
||||
#[derive(Deserialize)]
|
||||
struct Input {
|
||||
path: String,
|
||||
model_alias: String,
|
||||
provider_model: String,
|
||||
api_base: String,
|
||||
|
|
@ -21,7 +22,8 @@ async fn main() {
|
|||
Ok(input) => input,
|
||||
Err(error) => fail(error),
|
||||
};
|
||||
let result = litellm_ai_gateway::trace_parity::traced_messages_request(
|
||||
let result = litellm_ai_gateway::trace_parity::traced_request(
|
||||
input.path,
|
||||
input.model_alias,
|
||||
input.provider_model,
|
||||
input.api_base,
|
||||
|
|
|
|||
|
|
@ -29,14 +29,15 @@ pub struct TracedGatewayResponse {
|
|||
pub trace: Vec<litellm_core::observability::FunctionTraceEvent>,
|
||||
}
|
||||
|
||||
pub async fn traced_messages_request(
|
||||
pub async fn traced_request(
|
||||
path: String,
|
||||
model_alias: String,
|
||||
provider_model: String,
|
||||
api_base: String,
|
||||
body: Value,
|
||||
) -> TracedGatewayResponse {
|
||||
let trace = litellm_core::observability::FunctionTrace::default();
|
||||
let result = messages_request(model_alias, provider_model, api_base, body)
|
||||
let result = request(path, model_alias, provider_model, api_base, body)
|
||||
.with_subscriber(trace.dispatcher())
|
||||
.await;
|
||||
let events = trace.events();
|
||||
|
|
@ -54,7 +55,8 @@ pub async fn traced_messages_request(
|
|||
}
|
||||
}
|
||||
|
||||
pub async fn messages_request(
|
||||
pub async fn request(
|
||||
path: String,
|
||||
model_alias: String,
|
||||
provider_model: String,
|
||||
api_base: String,
|
||||
|
|
@ -75,7 +77,7 @@ pub async fn messages_request(
|
|||
};
|
||||
let request = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/v1/messages")
|
||||
.uri(path)
|
||||
.header(AUTHORIZATION, "Bearer trace-master-key")
|
||||
.header(CONTENT_TYPE, "application/json")
|
||||
.body(Body::from(body.to_string()))
|
||||
|
|
|
|||
|
|
@ -538,7 +538,7 @@ context_window_fallbacks: Optional[List] = None
|
|||
content_policy_fallbacks: Optional[List] = None
|
||||
allowed_fails: int = 3
|
||||
allow_dynamic_callback_disabling: bool = True
|
||||
num_retries_per_request: Optional[int] = None # for the request overall (incl. fallbacks + model retries)
|
||||
num_retries_per_request: Optional[int] = None # cap on Router retries of one model group; resets per fallback hop
|
||||
####### SECRET MANAGERS #####################
|
||||
secret_manager_client: Optional[Any] = (
|
||||
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import ast
|
||||
import contextvars
|
||||
import functools
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
|
|
@ -225,6 +226,35 @@ class AccessLogRedactionFilter(logging.Filter):
|
|||
_access_log_filter: Final = AccessLogRedactionFilter()
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _parse_disabled_access_log_paths(raw: str) -> frozenset[str]:
|
||||
return frozenset(stripped for path in raw.split(",") if (stripped := path.strip()))
|
||||
|
||||
|
||||
def _disabled_access_log_paths() -> frozenset[str]:
|
||||
"""Read the variable per record so a value loaded later via proxy config
|
||||
environment_variables or dotenv is honored."""
|
||||
return _parse_disabled_access_log_paths(os.getenv("LITELLM_DISABLE_ACCESS_LOG_PATHS", ""))
|
||||
|
||||
|
||||
class AccessLogPathFilter(logging.Filter):
|
||||
"""Drops uvicorn.access records for request paths listed in LITELLM_DISABLE_ACCESS_LOG_PATHS.
|
||||
|
||||
uvicorn passes record.args as (client_addr, method, full_path, http_version, status_code).
|
||||
"""
|
||||
|
||||
def filter(self, record: logging.LogRecord) -> bool:
|
||||
if not isinstance(record.args, tuple) or len(record.args) < 3:
|
||||
return True
|
||||
full_path: Final = record.args[2]
|
||||
if not isinstance(full_path, str):
|
||||
return True
|
||||
return full_path.partition("?")[0] not in _disabled_access_log_paths()
|
||||
|
||||
|
||||
_access_log_path_filter: Final = AccessLogPathFilter()
|
||||
|
||||
|
||||
def _get_max_string_length_stdout_log() -> int:
|
||||
"""Read the limit per record so a value loaded later via proxy config
|
||||
environment_variables is honored."""
|
||||
|
|
@ -663,6 +693,7 @@ def _redact_third_party_loggers() -> None:
|
|||
for name in _REDACTED_THIRD_PARTY_LOGGERS:
|
||||
logging.getLogger(name).addFilter(_secret_filter)
|
||||
for name in _REDACTED_ACCESS_LOGGERS:
|
||||
logging.getLogger(name).addFilter(_access_log_path_filter)
|
||||
logging.getLogger(name).addFilter(_access_log_filter)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -679,7 +679,14 @@ class Cache:
|
|||
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
|
||||
self.cache.set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Cache: exception in add_cache", e)
|
||||
self._log_add_cache_failure(e)
|
||||
|
||||
def _log_add_cache_failure(self, exc: Exception) -> None:
|
||||
message: Final = "LiteLLM Cache: exception in add_cache"
|
||||
if isinstance(self.cache, RedisCache):
|
||||
log_redis_failure(verbose_logger, logging.ERROR, message, exc)
|
||||
return
|
||||
verbose_logger.error("%s: %s", message, exc)
|
||||
|
||||
async def async_add_cache(self, result, dynamic_cache_object: BaseCache | None = None, **kwargs):
|
||||
"""
|
||||
|
|
@ -698,7 +705,7 @@ class Cache:
|
|||
else:
|
||||
await self.cache.async_set_cache(cache_key, cached_data, **kwargs)
|
||||
except Exception as e:
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Cache: exception in add_cache", e)
|
||||
self._log_add_cache_failure(e)
|
||||
|
||||
def _convert_to_cached_embedding(
|
||||
self,
|
||||
|
|
@ -877,7 +884,7 @@ class Cache:
|
|||
else:
|
||||
await self.cache.async_set_cache_pipeline(cache_list=cache_list, **kwargs)
|
||||
except Exception as e:
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Cache: exception in add_cache", e)
|
||||
self._log_add_cache_failure(e)
|
||||
|
||||
def should_use_cache(self, **kwargs):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import hashlib
|
|||
import inspect
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterator, Sequence
|
||||
from contextvars import ContextVar
|
||||
|
|
@ -32,6 +33,7 @@ from litellm.constants import (
|
|||
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
|
||||
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
|
||||
REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION,
|
||||
REDIS_TIMEOUT_LOG_INTERVAL,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
|
||||
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
|
||||
|
|
@ -340,7 +342,7 @@ def _explicit_causes(exc: BaseException) -> Iterator[BaseException]:
|
|||
current = current.__cause__
|
||||
|
||||
|
||||
def _is_redis_timeout_failure(exc: BaseException) -> bool:
|
||||
def is_redis_timeout_failure(exc: BaseException) -> bool:
|
||||
"""True when ``exc`` or any exception it was explicitly raised ``from`` is a timeout.
|
||||
|
||||
redis-py's blocking pool reports a pool wait timeout as ``ConnectionError`` chained from
|
||||
|
|
@ -414,7 +416,7 @@ def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseExcep
|
|||
"""
|
||||
if not _is_redis_health_failure(exc):
|
||||
return
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(exc))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(exc))
|
||||
_swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1)
|
||||
|
||||
|
||||
|
|
@ -422,13 +424,58 @@ class RedisCircuitBreakerOpenError(Exception):
|
|||
pass
|
||||
|
||||
|
||||
class _RedisTimeoutLogThrottle:
|
||||
"""Admits one Redis timeout log line per interval and counts the timeouts it suppressed in between."""
|
||||
|
||||
def __init__(self, interval: float, clock: Callable[[], float] = time.monotonic) -> None:
|
||||
self.interval = interval
|
||||
self._clock = clock
|
||||
self._lock = threading.Lock()
|
||||
self._last_logged_at: float | None = None
|
||||
self._suppressed = 0
|
||||
|
||||
def admit(self) -> int | None:
|
||||
"""Return the number of timeouts suppressed since the last admitted line, or None to suppress this one."""
|
||||
with self._lock:
|
||||
now: Final = self._clock()
|
||||
if self._last_logged_at is not None and now - self._last_logged_at < self.interval:
|
||||
self._suppressed += 1
|
||||
return None
|
||||
suppressed: Final = self._suppressed
|
||||
self._suppressed = 0
|
||||
self._last_logged_at = now
|
||||
return suppressed
|
||||
|
||||
|
||||
_redis_timeout_log_throttle: Final = _RedisTimeoutLogThrottle(REDIS_TIMEOUT_LOG_INTERVAL)
|
||||
|
||||
|
||||
def log_redis_failure(
|
||||
logger: logging.Logger, level: int, message: str, exc: BaseException, with_traceback: bool = False
|
||||
) -> None:
|
||||
if isinstance(exc, RedisCircuitBreakerOpenError):
|
||||
logger.debug("%s: %s", message, exc)
|
||||
logger.debug("%s: %s", message, exc, stacklevel=2)
|
||||
return
|
||||
logger.log(level, "%s: %s", message, exc, exc_info=exc if with_traceback else None)
|
||||
exc_info: Final = exc if with_traceback else None
|
||||
if not is_redis_timeout_failure(exc):
|
||||
logger.log(level, "%s: %s", message, exc, exc_info=exc_info, stacklevel=2)
|
||||
return
|
||||
suppressed: Final = _redis_timeout_log_throttle.admit()
|
||||
if suppressed is None:
|
||||
logger.debug("%s: %s", message, exc, stacklevel=2)
|
||||
return
|
||||
if suppressed == 0:
|
||||
logger.log(level, "%s: %s", message, exc, exc_info=exc_info, stacklevel=2)
|
||||
return
|
||||
logger.log(
|
||||
level,
|
||||
"%s: %s (%d more Redis timeouts since the previous Redis timeout line were logged at DEBUG)",
|
||||
message,
|
||||
exc,
|
||||
suppressed,
|
||||
exc_info=exc_info,
|
||||
stacklevel=2,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -475,7 +522,7 @@ async def _run_under_circuit_breaker(
|
|||
result: Final = await call()
|
||||
except Exception as e:
|
||||
if _is_redis_health_failure(e):
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(e))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(e))
|
||||
raise
|
||||
_exit_circuit_breaker(breaker, admission)
|
||||
return result
|
||||
|
|
@ -492,7 +539,7 @@ def _run_under_circuit_breaker_sync(
|
|||
result: Final = call()
|
||||
except Exception as e:
|
||||
if _is_redis_health_failure(e):
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(e))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(e))
|
||||
raise
|
||||
_exit_circuit_breaker(breaker, admission)
|
||||
return result
|
||||
|
|
@ -801,10 +848,8 @@ class RedisCache(BaseCache):
|
|||
## LOGGING ##
|
||||
end_time = time.time()
|
||||
_duration = end_time - start_time
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: increment_cache() - Got exception from REDIS %s, Writing value=%s",
|
||||
str(e),
|
||||
value,
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Redis Caching: increment_cache() - Got exception from REDIS", e
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -1010,11 +1055,8 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_set_cache <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async set() - Got exception from REDIS %s, key=%r, value=%r",
|
||||
str(e),
|
||||
key,
|
||||
value,
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Redis Caching: async set() - Got exception from REDIS", e
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -1062,10 +1104,8 @@ class RedisCache(BaseCache):
|
|||
event_metadata={"key": key},
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async set() - Got exception from REDIS %s, Writing value=%s",
|
||||
str(e),
|
||||
value,
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Redis Caching: async set() - Got exception from REDIS", e
|
||||
)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
|
||||
|
|
@ -1112,7 +1152,6 @@ class RedisCache(BaseCache):
|
|||
start_time: Final = time.time()
|
||||
|
||||
print_verbose(f"Set Async Redis Cache: key list: {cache_list}\nttl={ttl}, redis_version={self.redis_version}")
|
||||
cache_value: Final = None
|
||||
try:
|
||||
async with _redis_client.pipeline(transaction=False) as pipe:
|
||||
results: Final = await self._pipeline_helper(pipe, cache_list, ttl)
|
||||
|
|
@ -1149,10 +1188,11 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
)
|
||||
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async set_cache_pipeline() - Got exception from REDIS %s, Writing value=%s",
|
||||
str(e),
|
||||
cache_value,
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Redis Caching: async set_cache_pipeline() - Got exception from REDIS",
|
||||
e,
|
||||
)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
|
||||
|
|
@ -1191,8 +1231,11 @@ class RedisCache(BaseCache):
|
|||
end_time=time.time(),
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async_set_cache_pipeline_with_ttls() - Got exception from REDIS %s", str(e)
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Redis Caching: async_set_cache_pipeline_with_ttls() - Got exception from REDIS",
|
||||
e,
|
||||
)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
|
||||
|
|
@ -1235,10 +1278,8 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
)
|
||||
# NON blocking - notify users Redis is throwing an exception
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async set() - Got exception from REDIS %s, Writing value=%s",
|
||||
str(e),
|
||||
value,
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "LiteLLM Redis Caching: async set() - Got exception from REDIS", e
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -1274,10 +1315,11 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
)
|
||||
# NON blocking - notify users Redis is throwing an exception
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async set_cache_sadd() - Got exception from REDIS %s, Writing value=%s",
|
||||
str(e),
|
||||
value,
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Redis Caching: async set_cache_sadd() - Got exception from REDIS",
|
||||
e,
|
||||
)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
|
||||
|
|
@ -1359,10 +1401,11 @@ class RedisCache(BaseCache):
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async async_increment() - Got exception from REDIS %s, Writing value=%s",
|
||||
str(e),
|
||||
value,
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Redis Caching: async async_increment() - Got exception from REDIS",
|
||||
e,
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -1448,7 +1491,9 @@ class RedisCache(BaseCache):
|
|||
print_verbose(f"Got Redis Cache: key: {key}, cached_response {cached_response}")
|
||||
return self._get_cache_logic(cached_response=cached_response)
|
||||
except Exception as e:
|
||||
verbose_logger.error("litellm.caching.caching: get() - Got exception from REDIS: %s", e)
|
||||
log_redis_failure(
|
||||
verbose_logger, logging.ERROR, "litellm.caching.caching: get() - Got exception from REDIS", e
|
||||
)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
|
||||
def _run_redis_mget_operation(self, keys: list[str]) -> Sequence[bytes | str | None]:
|
||||
|
|
@ -1526,7 +1571,7 @@ class RedisCache(BaseCache):
|
|||
end_time=failed_at,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
verbose_logger.error("Error occurred in batch get cache - %s", e)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "Error occurred in batch get cache", e)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
return key_value_dict
|
||||
|
||||
|
|
@ -1645,7 +1690,7 @@ class RedisCache(BaseCache):
|
|||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
)
|
||||
verbose_logger.error("Error occurred in async batch get cache - %s", e)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "Error occurred in async batch get cache", e)
|
||||
_record_swallowed_redis_failure(self._circuit_breaker, e)
|
||||
return key_value_dict
|
||||
|
||||
|
|
@ -1870,9 +1915,11 @@ class RedisCache(BaseCache):
|
|||
parent_otel_span=_get_parent_otel_span_from_kwargs(kwargs),
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async increment_pipeline() - Got exception from REDIS %s",
|
||||
str(e),
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Redis Caching: async increment_pipeline() - Got exception from REDIS",
|
||||
e,
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -1949,7 +1996,7 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_rpush <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error("LiteLLM Redis Cache RPUSH: - Got exception from REDIS : %s", e)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Redis Cache RPUSH: - Got exception from REDIS", e)
|
||||
raise e
|
||||
|
||||
async def _pipeline_rpush_helper(
|
||||
|
|
@ -2017,9 +2064,11 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_rpush_pipeline <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async_rpush_pipeline() - Got exception from REDIS %s",
|
||||
str(e),
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Redis Caching: async_rpush_pipeline() - Got exception from REDIS",
|
||||
e,
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -2095,7 +2144,7 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_lpop <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error("LiteLLM Redis Cache LPOP: - Got exception from REDIS : %s", e)
|
||||
log_redis_failure(verbose_logger, logging.ERROR, "LiteLLM Redis Cache LPOP: - Got exception from REDIS", e)
|
||||
raise e
|
||||
|
||||
async def _pipeline_lpop_helper(
|
||||
|
|
@ -2206,8 +2255,10 @@ class RedisCache(BaseCache):
|
|||
call_type=f"async_lpop_pipeline <- {_get_call_stack_info()}",
|
||||
)
|
||||
)
|
||||
verbose_logger.error(
|
||||
"LiteLLM Redis Caching: async_lpop_pipeline() - Got exception from REDIS %s",
|
||||
str(e),
|
||||
log_redis_failure(
|
||||
verbose_logger,
|
||||
logging.ERROR,
|
||||
"LiteLLM Redis Caching: async_lpop_pipeline() - Got exception from REDIS",
|
||||
e,
|
||||
)
|
||||
raise e
|
||||
|
|
|
|||
|
|
@ -227,6 +227,9 @@ PRE_CALL_EXECUTED_GUARDRAILS_KEY: Final = "_pre_call_executed_guardrails"
|
|||
# Attribute stamped on log_guardrail_information wrappers so __init_subclass__ does not wrap them again
|
||||
LOGS_GUARDRAIL_INFORMATION_MARKER: Final = "_litellm_logs_guardrail_information"
|
||||
|
||||
# llm_provider stamped on proxy-side rate limit errors when the model resolves to no deployment
|
||||
PROXY_LLM_PROVIDER_FALLBACK: Final = "litellm_proxy"
|
||||
|
||||
# Generic fallback for unknown models
|
||||
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET: Final = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128)
|
||||
|
|
@ -311,6 +314,12 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float(
|
|||
# RFC 6455 caps the close frame payload at 125 bytes, 2 of which carry the status code
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
|
||||
|
||||
BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update"
|
||||
BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY: Final = "litellm.bedrock_realtime.session_committed"
|
||||
BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY: Final = "litellm.bedrock_realtime.committed_failure"
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY: Final = "realtime_session_success_logged"
|
||||
REALTIME_SESSION_FAILURE_LOGGED_KEY: Final = "realtime_session_failure_logged"
|
||||
|
||||
# SSL/TLS cipher configuration for faster handshakes
|
||||
# Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones
|
||||
# This balances performance with broad compatibility
|
||||
|
|
@ -461,6 +470,7 @@ REDIS_CIRCUIT_BREAKER_ENABLED: Final = os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED"
|
|||
# minimum seconds a timeout-only failure streak must span before it can open the breaker,
|
||||
# so one event-loop stall timing out many queued calls at once does not trip it
|
||||
REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION: Final = float(os.getenv("REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION", 5.0))
|
||||
REDIS_TIMEOUT_LOG_INTERVAL: Final = float(os.getenv("REDIS_TIMEOUT_LOG_INTERVAL", "5.0"))
|
||||
# Seconds of idle before a Redis cluster connection is validated with a PING and
|
||||
# reconnected if dead, so a connection silently dropped by a cluster restart
|
||||
# (e.g. ElastiCache Serverless maintenance) is not reused while broken
|
||||
|
|
|
|||
|
|
@ -97,6 +97,7 @@ from litellm.llms.vertex_ai.cost_calculator import cost_router as google_cost_ro
|
|||
from litellm.llms.xai.cost_calculator import cost_per_token as xai_cost_per_token
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
from litellm.types.llms.base import CachedTokensDetails
|
||||
from litellm.types.llms.openai import (
|
||||
HttpxBinaryResponseContent,
|
||||
ImageGenerationRequestQuality,
|
||||
|
|
@ -2381,6 +2382,46 @@ def _summable_prompt_token_fields(prompt_tokens_details: BaseModel) -> list[str]
|
|||
return [attr for attr in field_names if attr != "cache_creation_tokens"]
|
||||
|
||||
|
||||
def _combine_cached_tokens_details(
|
||||
current: CachedTokensDetails | None, new: CachedTokensDetails
|
||||
) -> CachedTokensDetails:
|
||||
def _sum_optional(current_value: int | None, new_value: int | None) -> int | None:
|
||||
if current_value is None and new_value is None:
|
||||
return None
|
||||
return (current_value or 0) + (new_value or 0)
|
||||
|
||||
return CachedTokensDetails(
|
||||
text_tokens=_sum_optional(current.text_tokens if current is not None else None, new.text_tokens),
|
||||
audio_tokens=_sum_optional(current.audio_tokens if current is not None else None, new.audio_tokens),
|
||||
image_tokens=_sum_optional(current.image_tokens if current is not None else None, new.image_tokens),
|
||||
)
|
||||
|
||||
|
||||
def _combine_prompt_tokens_details(
|
||||
current: PromptTokensDetailsWrapper | None, new: PromptTokensDetailsWrapper
|
||||
) -> PromptTokensDetailsWrapper:
|
||||
base: Final = current if current is not None else PromptTokensDetailsWrapper()
|
||||
base_values: Final = MappingProxyType(
|
||||
{attr: getattr(base, attr) for attr in type(base).model_fields if hasattr(base, attr)}
|
||||
)
|
||||
summed: Final = MappingProxyType(
|
||||
{
|
||||
attr: (getattr(base, attr, 0) or 0) + (getattr(new, attr) or 0)
|
||||
for attr in _summable_prompt_token_fields(new)
|
||||
if hasattr(new, attr) and isinstance(getattr(new, attr) or 0, (int, float))
|
||||
}
|
||||
)
|
||||
new_cached_tokens_details: Final = getattr(new, "cached_tokens_details", None)
|
||||
cached_tokens_details: Final = (
|
||||
_combine_cached_tokens_details(getattr(base, "cached_tokens_details", None), new_cached_tokens_details)
|
||||
if isinstance(new_cached_tokens_details, CachedTokensDetails)
|
||||
else getattr(base, "cached_tokens_details", None)
|
||||
)
|
||||
return PromptTokensDetailsWrapper(
|
||||
**MappingProxyType({**base_values, **summed, "cached_tokens_details": cached_tokens_details})
|
||||
)
|
||||
|
||||
|
||||
class BaseTokenUsageProcessor:
|
||||
@staticmethod
|
||||
def combine_usage_objects(usage_objects: list[Usage]) -> Usage:
|
||||
|
|
@ -2389,7 +2430,6 @@ class BaseTokenUsageProcessor:
|
|||
"""
|
||||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
|
@ -2408,27 +2448,10 @@ class BaseTokenUsageProcessor:
|
|||
and isinstance(current_val, (int, float))
|
||||
):
|
||||
setattr(combined, attr, current_val + new_val)
|
||||
# Handle nested prompt_tokens_details
|
||||
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
|
||||
if not hasattr(combined, "prompt_tokens_details") or not combined.prompt_tokens_details:
|
||||
combined.prompt_tokens_details = PromptTokensDetailsWrapper()
|
||||
|
||||
# Check what keys exist in the model's prompt_tokens_details
|
||||
# Access model_fields on the class, not the instance, to avoid Pydantic 2.11+ deprecation warnings
|
||||
for attr in _summable_prompt_token_fields(usage.prompt_tokens_details):
|
||||
if (
|
||||
hasattr(usage.prompt_tokens_details, attr)
|
||||
and not attr.startswith("_")
|
||||
and not callable(_attribute_value(usage.prompt_tokens_details, attr))
|
||||
):
|
||||
current_val = getattr(combined.prompt_tokens_details, attr, 0) or 0
|
||||
new_val = getattr(usage.prompt_tokens_details, attr, 0) or 0
|
||||
if new_val is not None and isinstance(new_val, (int, float)):
|
||||
setattr(
|
||||
combined.prompt_tokens_details,
|
||||
attr,
|
||||
current_val + new_val,
|
||||
)
|
||||
combined.prompt_tokens_details = _combine_prompt_tokens_details(
|
||||
getattr(combined, "prompt_tokens_details", None), usage.prompt_tokens_details
|
||||
)
|
||||
|
||||
# Handle nested completion_tokens_details
|
||||
if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from typing import Final
|
|||
|
||||
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData
|
||||
from litellm.integrations.otel.mappers.utils import (
|
||||
MAX_MESSAGE_ATTRS_PER_SPAN,
|
||||
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
|
||||
collect,
|
||||
drop_none,
|
||||
|
|
@ -26,6 +27,8 @@ from litellm.integrations.otel.model.payloads import (
|
|||
ToolDefinition,
|
||||
)
|
||||
|
||||
_MAX_INDEXED_MESSAGES: Final = MAX_MESSAGE_ATTRS_PER_SPAN // 2
|
||||
|
||||
|
||||
class OpenInferenceMapper:
|
||||
"""Emits OpenInference attributes for LLM_CALL spans.
|
||||
|
|
@ -84,27 +87,44 @@ class OpenInferenceMapper:
|
|||
return {}
|
||||
|
||||
def _llm_call(self, data: LLMCallSpanData) -> AttributeMap:
|
||||
outputs: Final = output_messages(data)
|
||||
indexed_in, indexed_out = self._indexed_split(len(data.messages_in), len(outputs))
|
||||
return {
|
||||
**collect(self._LLM_CALL_ATTRS, data),
|
||||
**collect(self._BLOB_ATTRS, data),
|
||||
**self._messages("llm.input_messages", "input.value", data.messages_in),
|
||||
**self._messages("llm.output_messages", "output.value", output_messages(data)),
|
||||
**self._messages(
|
||||
"llm.input_messages",
|
||||
"input.value",
|
||||
data.messages_in,
|
||||
self._prompt_positions(len(data.messages_in), indexed_in),
|
||||
),
|
||||
**self._messages("llm.output_messages", "output.value", outputs, range(indexed_out)),
|
||||
**self._tools(data),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _messages(prefix: str, value_key: str, messages: Sequence[object]) -> AttributeMap:
|
||||
"""Per-message ``{prefix}.{idx}.message.*`` keys + the ``value_key`` blob."""
|
||||
def _indexed_split(inputs: int, outputs: int) -> tuple[int, int]:
|
||||
"""Prompt and response share one allowance; the response is reserved at least half of it."""
|
||||
indexed_out: Final = min(outputs, max(_MAX_INDEXED_MESSAGES // 2, _MAX_INDEXED_MESSAGES - inputs))
|
||||
return _MAX_INDEXED_MESSAGES - indexed_out, indexed_out
|
||||
|
||||
@staticmethod
|
||||
def _prompt_positions(total: int, indexed: int) -> tuple[int, ...]:
|
||||
"""Prompt messages that get per-index attributes: message 0 and the most recent turns."""
|
||||
if total <= indexed:
|
||||
return tuple(range(total))
|
||||
return (0, *range(total - indexed + 1, total))
|
||||
|
||||
@staticmethod
|
||||
def _messages(prefix: str, value_key: str, messages: Sequence[object], positions: Sequence[int]) -> AttributeMap:
|
||||
"""``{prefix}.{idx}.message.*`` keys for the messages at ``positions`` + the ``value_key`` blob of all."""
|
||||
parsed: Final = [(m.get("role") if isinstance(m, dict) else None, message_content(m)) for m in messages]
|
||||
attrs: Final = drop_none(
|
||||
{
|
||||
key: value
|
||||
for idx, (role, content) in enumerate(parsed)
|
||||
for idx, (role, content) in ((idx, parsed[idx]) for idx in positions)
|
||||
for key, value in (
|
||||
(
|
||||
f"{prefix}.{idx}.message.role",
|
||||
role if isinstance(role, str) else None,
|
||||
),
|
||||
(f"{prefix}.{idx}.message.role", role if isinstance(role, str) else None),
|
||||
(f"{prefix}.{idx}.message.content", content),
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,6 +32,14 @@ core telemetry no matter how many vocabularies are configured.
|
|||
"""
|
||||
|
||||
|
||||
MAX_MESSAGE_ATTRS_PER_SPAN: Final = DEFAULT_SPAN_ATTRIBUTE_LIMIT // 8
|
||||
"""Span-wide ceiling on per-index chat message attributes, prompt and response together.
|
||||
|
||||
An eighth is the largest share that still fits beside the tool ceiling and the core
|
||||
of every vocabulary at once. The complete conversation still rides the JSON blobs.
|
||||
"""
|
||||
|
||||
|
||||
def tool_attr_budget(vocabularies: int) -> int:
|
||||
"""Split the span-wide tool-definition ceiling across active vocabularies."""
|
||||
return MAX_TOOL_DEFINITION_ATTRS_PER_SPAN // max(vocabularies, 1)
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from pydantic import BaseModel
|
|||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import PROXY_LLM_PROVIDER_FALLBACK
|
||||
from litellm.exceptions import (
|
||||
validate_rate_limit_category,
|
||||
validate_rate_limit_type,
|
||||
|
|
@ -2581,6 +2582,15 @@ class PrometheusLogger(CustomLogger):
|
|||
)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _extract_api_provider_from_exception(exception: Exception) -> str | None:
|
||||
if not isinstance(exception, litellm.exceptions.RateLimitError):
|
||||
return None
|
||||
llm_provider: Final = exception.llm_provider
|
||||
if not llm_provider or llm_provider == PROXY_LLM_PROVIDER_FALLBACK:
|
||||
return None
|
||||
return llm_provider
|
||||
|
||||
async def async_post_call_failure_hook(
|
||||
self,
|
||||
request_data: dict,
|
||||
|
|
@ -2616,7 +2626,9 @@ class PrometheusLogger(CustomLogger):
|
|||
_metadata: Final = request_data.get("metadata", {}) or {}
|
||||
model_id: Final = _metadata.get("model_info", {}).get("id") or request_data.get("model_info", {}).get("id")
|
||||
rate_limit_category, rate_limit_type = self._extract_rate_limit_labels(original_exception)
|
||||
api_provider: Final = self._extract_api_provider_from_request_data(request_data)
|
||||
api_provider: Final = self._extract_api_provider_from_request_data(
|
||||
request_data
|
||||
) or self._extract_api_provider_from_exception(original_exception)
|
||||
enum_values: Final = UserAPIKeyLabelValues(
|
||||
end_user=user_api_key_dict.end_user_id,
|
||||
user=user_api_key_dict.user_id,
|
||||
|
|
|
|||
|
|
@ -303,6 +303,16 @@ def get_metadata_variable_name_from_kwargs(
|
|||
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
|
||||
|
||||
def max_retries_per_request_hit(kwargs: Mapping[str, object], num_retries_per_request: int | None) -> bool:
|
||||
if num_retries_per_request is None:
|
||||
return False
|
||||
metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
|
||||
if not isinstance(metadata, Mapping):
|
||||
return False
|
||||
attempted_retries: Final = metadata.get("attempted_retries")
|
||||
return type(attempted_retries) is int and 0 < attempted_retries and num_retries_per_request <= attempted_retries
|
||||
|
||||
|
||||
def get_or_create_metadata_bucket(
|
||||
request_data: dict,
|
||||
) -> tuple[Literal["metadata", "litellm_metadata"], dict]:
|
||||
|
|
|
|||
|
|
@ -34,6 +34,11 @@ rules never mix the two and never use ``extends``. A rule whose
|
|||
Rules are only consulted after exact and case-insensitive lookups miss, so an
|
||||
exact cost-map entry always takes precedence over any rule.
|
||||
|
||||
Rules flagged with ``fill_missing_for_providers: [..]`` also fill only keys
|
||||
missing from an exact cost-map entry when the entry's ``litellm_provider`` is
|
||||
listed, while values already present on the entry win on conflict. Only flagged
|
||||
capability rules participate in this fill; routing rules never do.
|
||||
|
||||
Patterns are matched case-insensitively with ``re.search`` and are not implicitly
|
||||
anchored: a rule must include ``^`` and ``$`` to bind to the whole model name,
|
||||
otherwise it matches as a substring. Keeping anchoring in the regex makes the rule
|
||||
|
|
@ -46,17 +51,19 @@ Rules are compiled and classified once, at install time. The match functions are
|
|||
O(number of rules); callers must only invoke them on a cache miss.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
verbose_logger: Final = logging.getLogger("LiteLLM")
|
||||
NAME_FIELD: Final = "name"
|
||||
PATTERN_FIELD: Final = "pattern"
|
||||
MODEL_INFO_FIELD: Final = "model_info"
|
||||
PROVIDER_KEY: Final = "litellm_provider"
|
||||
LEGACY_EXTENDS_FIELD: Final = "extends"
|
||||
FILL_MISSING_FOR_PROVIDERS_FIELD: Final = "fill_missing_for_providers"
|
||||
|
||||
|
||||
def _resolve_legacy_extends(rules: list) -> list:
|
||||
|
|
@ -98,11 +105,28 @@ class _RoutingRule:
|
|||
class _CapabilityRule:
|
||||
pattern: re.Pattern
|
||||
model_info: dict
|
||||
fill_missing_for_providers: frozenset[str]
|
||||
|
||||
|
||||
_CompiledRule = _RoutingRule | _CapabilityRule
|
||||
|
||||
|
||||
def _parse_fill_missing_for_providers(rule: Mapping[str, object], pattern_label: object) -> frozenset[str] | None:
|
||||
if FILL_MISSING_FOR_PROVIDERS_FIELD not in rule:
|
||||
return frozenset()
|
||||
raw_fill_missing_for_providers: Final = rule.get(FILL_MISSING_FOR_PROVIDERS_FIELD)
|
||||
if not isinstance(raw_fill_missing_for_providers, (list, tuple)) or not all(
|
||||
isinstance(provider, str) for provider in raw_fill_missing_for_providers
|
||||
):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: skipping malformed fallback generalization rule %s ('%s' must be a list of provider strings).",
|
||||
rule.get(NAME_FIELD, pattern_label),
|
||||
FILL_MISSING_FOR_PROVIDERS_FIELD,
|
||||
)
|
||||
return None
|
||||
return frozenset(raw_fill_missing_for_providers)
|
||||
|
||||
|
||||
def _compile_rule(rule: object) -> tuple[_CompiledRule, ...]:
|
||||
if not isinstance(rule, dict):
|
||||
return ()
|
||||
|
|
@ -125,8 +149,17 @@ def _compile_rule(rule: object) -> tuple[_CompiledRule, ...]:
|
|||
e,
|
||||
)
|
||||
return ()
|
||||
fill_missing_for_providers: Final = _parse_fill_missing_for_providers(rule, pattern)
|
||||
if fill_missing_for_providers is None:
|
||||
return ()
|
||||
if PROVIDER_KEY not in model_info:
|
||||
return (_CapabilityRule(pattern=compiled, model_info=model_info),)
|
||||
return (
|
||||
_CapabilityRule(
|
||||
pattern=compiled,
|
||||
model_info=model_info,
|
||||
fill_missing_for_providers=fill_missing_for_providers,
|
||||
),
|
||||
)
|
||||
provider: Final = model_info[PROVIDER_KEY]
|
||||
if not isinstance(provider, str):
|
||||
verbose_logger.warning(
|
||||
|
|
@ -140,7 +173,11 @@ def _compile_rule(rule: object) -> tuple[_CompiledRule, ...]:
|
|||
return (_RoutingRule(pattern=compiled, provider=provider),)
|
||||
return (
|
||||
_RoutingRule(pattern=compiled, provider=provider),
|
||||
_CapabilityRule(pattern=compiled, model_info=model_info),
|
||||
_CapabilityRule(
|
||||
pattern=compiled,
|
||||
model_info=model_info,
|
||||
fill_missing_for_providers=fill_missing_for_providers,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -151,6 +188,7 @@ class _FallbackGeneralizations:
|
|||
self.rules: list = []
|
||||
self.routing_rules: tuple = ()
|
||||
self.capability_rules: tuple = ()
|
||||
self.fill_missing_rules: tuple[_CapabilityRule, ...] = ()
|
||||
|
||||
def set_rules(self, rules: list | None) -> None:
|
||||
installed: Final = rules if isinstance(rules, list) else []
|
||||
|
|
@ -158,6 +196,7 @@ class _FallbackGeneralizations:
|
|||
self.rules = installed
|
||||
self.routing_rules = tuple(rule for rule in compiled if isinstance(rule, _RoutingRule))
|
||||
self.capability_rules = tuple(rule for rule in compiled if isinstance(rule, _CapabilityRule))
|
||||
self.fill_missing_rules = tuple(rule for rule in self.capability_rules if rule.fill_missing_for_providers)
|
||||
|
||||
def match_routing(self, model: str) -> str | None:
|
||||
if not model:
|
||||
|
|
@ -175,6 +214,21 @@ class _FallbackGeneralizations:
|
|||
return None
|
||||
return {key: value for model_info in matched for key, value in model_info.items()}
|
||||
|
||||
def match_fill_missing(self, model: str, provider: str) -> Mapping[str, object] | None:
|
||||
if not model or not provider:
|
||||
return None
|
||||
matched = tuple(
|
||||
rule.model_info
|
||||
for rule in self.fill_missing_rules
|
||||
if provider in rule.fill_missing_for_providers and rule.pattern.search(model) is not None
|
||||
)
|
||||
if not matched:
|
||||
return None
|
||||
fill_missing: Final[Mapping[str, object]] = {
|
||||
key: value for model_info in matched for key, value in model_info.items() if key != PROVIDER_KEY
|
||||
}
|
||||
return fill_missing or None
|
||||
|
||||
|
||||
_registry: Final = _FallbackGeneralizations()
|
||||
|
||||
|
|
@ -210,3 +264,14 @@ def match_capability_generalizations(model: str) -> dict | None:
|
|||
capability rule matches. O(number of rules); only call once exact lookups have missed.
|
||||
"""
|
||||
return _registry.match_capabilities(model)
|
||||
|
||||
|
||||
def match_fill_missing_generalizations(model: str, provider: str) -> Mapping[str, object] | None:
|
||||
"""Return flagged capability rules matching ``model`` for ``provider``.
|
||||
|
||||
Later rules override earlier ones on key conflicts. Only rules listing
|
||||
``provider`` in ``fill_missing_for_providers`` contribute. Returns ``None``
|
||||
when no flagged rule matches. O(number of rules); only call once exact
|
||||
lookups have matched.
|
||||
"""
|
||||
return _registry.match_fill_missing(model, provider)
|
||||
|
|
|
|||
|
|
@ -556,7 +556,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
# ids leaking into a different, later request on the same thread. Sync
|
||||
# support is deferred to a follow-up PR with its own safe-restore
|
||||
# mechanism; async calls (the proxy's only call path) are unaffected.
|
||||
if supports_correlation_logging:
|
||||
if supports_correlation_logging and litellm.request_correlation_in_logs:
|
||||
set_trace_id(self.litellm_trace_id)
|
||||
set_session_id(self.litellm_session_id)
|
||||
# set_trace_id()/set_session_id() sanitize (strip control chars, bound
|
||||
|
|
@ -2442,7 +2442,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
call) would leave the outer request's subsequent log lines stamped with
|
||||
the nested call's trace_id/session_id instead of its own.
|
||||
|
||||
Uses a plain set() of the captured pre-call value rather than
|
||||
Uses a plain contextvar set() of the captured pre-call value rather than
|
||||
contextvars.Token-based reset(), since this can end up called from a
|
||||
different asyncio Task/context than __init__ ran in (e.g. the request
|
||||
task's own wrapper() finally block, plus async_success_handler
|
||||
|
|
@ -2453,8 +2453,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
that Task's view of the contextvars, so calling it multiple times
|
||||
(once per Task involved in this attempt) is required, not just safe.
|
||||
"""
|
||||
set_trace_id(self._pre_call_trace_id)
|
||||
set_session_id(self._pre_call_session_id)
|
||||
trace_id_var.set(self._pre_call_trace_id)
|
||||
session_id_var.set(self._pre_call_session_id)
|
||||
|
||||
def _restore_correlation_context_if_unclaimed(self) -> None:
|
||||
"""Guarded variant for __del__-triggered cleanup only.
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ from types import MappingProxyType
|
|||
from typing import Any, Final, Literal, TypedDict, cast
|
||||
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_billing_time
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -772,6 +774,7 @@ def calculate_cache_writing_cost(
|
|||
|
||||
class PromptTokensDetailsResult(TypedDict):
|
||||
cache_hit_tokens: int
|
||||
cache_hit_audio_tokens: ReadOnly[int]
|
||||
cache_creation_tokens: int
|
||||
cache_creation_token_details: CacheCreationTokenDetails | None
|
||||
text_tokens: int
|
||||
|
|
@ -802,12 +805,34 @@ def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
|
|||
)
|
||||
or None
|
||||
)
|
||||
text_tokens: Final = (
|
||||
cast(int | None, getattr(usage.prompt_tokens_details, "text_tokens", None))
|
||||
or 0 # default to prompt tokens, if this field is not set
|
||||
cached_tokens_details: Final = getattr(usage.prompt_tokens_details, "cached_tokens_details", None)
|
||||
cached_audio_tokens: Final = min(
|
||||
_get_token_detail_value(cached_tokens_details, "audio_tokens") or 0, cache_hit_tokens
|
||||
)
|
||||
cached_text_tokens: Final = min(
|
||||
_get_token_detail_value(cached_tokens_details, "text_tokens") or 0,
|
||||
cache_hit_tokens - cached_audio_tokens,
|
||||
)
|
||||
cached_image_tokens: Final = min(
|
||||
_get_token_detail_value(cached_tokens_details, "image_tokens") or 0,
|
||||
cache_hit_tokens - cached_audio_tokens - cached_text_tokens,
|
||||
)
|
||||
text_tokens: Final = max(
|
||||
(
|
||||
cast(int | None, getattr(usage.prompt_tokens_details, "text_tokens", None))
|
||||
or 0 # default to prompt tokens, if this field is not set
|
||||
)
|
||||
- cached_text_tokens,
|
||||
0,
|
||||
)
|
||||
audio_tokens: Final = max(
|
||||
(cast(int | None, getattr(usage.prompt_tokens_details, "audio_tokens", 0)) or 0) - cached_audio_tokens,
|
||||
0,
|
||||
)
|
||||
image_tokens: Final = max(
|
||||
(cast(int | None, getattr(usage.prompt_tokens_details, "image_tokens", 0)) or 0) - cached_image_tokens,
|
||||
0,
|
||||
)
|
||||
audio_tokens: Final = cast(int | None, getattr(usage.prompt_tokens_details, "audio_tokens", 0)) or 0
|
||||
image_tokens: Final = cast(int | None, getattr(usage.prompt_tokens_details, "image_tokens", 0)) or 0
|
||||
video_tokens: Final = _coerce_token_count(getattr(usage.prompt_tokens_details, "video_tokens", 0))
|
||||
character_count: Final = (
|
||||
cast(
|
||||
|
|
@ -835,6 +860,7 @@ def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
|
|||
|
||||
return PromptTokensDetailsResult(
|
||||
cache_hit_tokens=cache_hit_tokens,
|
||||
cache_hit_audio_tokens=cached_audio_tokens,
|
||||
cache_creation_tokens=cache_creation_tokens,
|
||||
cache_creation_token_details=cache_creation_token_details,
|
||||
text_tokens=text_tokens,
|
||||
|
|
@ -918,7 +944,16 @@ def _calculate_input_cost(
|
|||
prompt_cost = float(prompt_tokens_details["text_tokens"]) * prompt_base_cost
|
||||
|
||||
### CACHE READ COST - Now uses tiered pricing
|
||||
prompt_cost += float(prompt_tokens_details["cache_hit_tokens"]) * cache_read_cost
|
||||
cache_hit_audio_tokens: Final = prompt_tokens_details["cache_hit_audio_tokens"]
|
||||
audio_cache_read_rate: Final = _get_cost_per_unit(
|
||||
model_info,
|
||||
_get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier),
|
||||
None,
|
||||
)
|
||||
prompt_cost += float(prompt_tokens_details["cache_hit_tokens"] - cache_hit_audio_tokens) * cache_read_cost
|
||||
prompt_cost += float(cache_hit_audio_tokens) * (
|
||||
audio_cache_read_rate if audio_cache_read_rate is not None else cache_read_cost
|
||||
)
|
||||
|
||||
### AUDIO COST
|
||||
if prompt_tokens_details["audio_tokens"]:
|
||||
|
|
@ -1149,6 +1184,7 @@ def generic_cost_per_token(
|
|||
### PROCESSING COST
|
||||
prompt_tokens_details = PromptTokensDetailsResult(
|
||||
cache_hit_tokens=0,
|
||||
cache_hit_audio_tokens=0,
|
||||
cache_creation_tokens=0,
|
||||
cache_creation_token_details=None,
|
||||
text_tokens=usage.prompt_tokens,
|
||||
|
|
@ -1319,6 +1355,7 @@ class BilledTokenRates:
|
|||
input_cost_per_token: float
|
||||
output_cost_per_token: float
|
||||
cache_read_input_token_cost: float
|
||||
cache_read_input_audio_token_cost: float
|
||||
cache_creation_input_token_cost: float
|
||||
cache_creation_input_token_cost_above_1hr: float
|
||||
output_cost_per_reasoning_token: float
|
||||
|
|
@ -1330,6 +1367,7 @@ class BilledTokenRates:
|
|||
input_cost_per_token=self.input_cost_per_token * multiplier,
|
||||
output_cost_per_token=self.output_cost_per_token * multiplier,
|
||||
cache_read_input_token_cost=self.cache_read_input_token_cost * multiplier,
|
||||
cache_read_input_audio_token_cost=self.cache_read_input_audio_token_cost * multiplier,
|
||||
cache_creation_input_token_cost=self.cache_creation_input_token_cost * multiplier,
|
||||
cache_creation_input_token_cost_above_1hr=self.cache_creation_input_token_cost_above_1hr * multiplier,
|
||||
output_cost_per_reasoning_token=self.output_cost_per_reasoning_token * multiplier,
|
||||
|
|
@ -1353,15 +1391,16 @@ def _reasoning_token_count(usage: Usage) -> int:
|
|||
return parsed or _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
|
||||
|
||||
|
||||
def _cache_token_counts(usage: Usage) -> tuple[int, int, CacheCreationTokenDetails | None]:
|
||||
"""(cache read tokens, cache creation tokens, cache creation details): read from prompt_tokens_details
|
||||
first, then the private top-level counters the Usage constructor mirrors cache tokens onto for
|
||||
providers/callers that bypass the details."""
|
||||
def _cache_token_counts(usage: Usage) -> tuple[int, int, int, CacheCreationTokenDetails | None]:
|
||||
"""(cache read tokens, cached audio tokens, cache creation tokens, cache creation details): read from
|
||||
prompt_tokens_details first, then the private top-level counters the Usage constructor mirrors cache
|
||||
tokens onto for providers/callers that bypass the details."""
|
||||
parsed: Final = parse_prompt_tokens_details(usage) if usage.prompt_tokens_details is not None else None
|
||||
parsed_read: Final = parsed["cache_hit_tokens"] if parsed is not None else 0
|
||||
parsed_creation: Final = parsed["cache_creation_tokens"] if parsed is not None else 0
|
||||
return (
|
||||
parsed_read or _coerce_token_count(getattr(usage, "_cache_read_input_tokens", 0)),
|
||||
parsed["cache_hit_audio_tokens"] if parsed is not None else 0,
|
||||
parsed_creation or _coerce_token_count(getattr(usage, "_cache_creation_input_tokens", 0)),
|
||||
parsed["cache_creation_token_details"] if parsed is not None else None,
|
||||
)
|
||||
|
|
@ -1372,11 +1411,13 @@ def _custom_pricing_rates(custom_cost_per_token: CostPerToken) -> BilledTokenRat
|
|||
cache rates (else the input rate) and reasoning at the output rate, as _cost_per_token_custom_pricing_helper does."""
|
||||
input_rate: Final = custom_cost_per_token["input_cost_per_token"]
|
||||
output_rate: Final = custom_cost_per_token["output_cost_per_token"]
|
||||
cache_read_rate: Final = custom_cost_per_token.get("cache_read_input_token_cost", input_rate)
|
||||
cache_creation_rate: Final = custom_cost_per_token.get("cache_creation_input_token_cost", input_rate)
|
||||
return BilledTokenRates(
|
||||
input_cost_per_token=input_rate,
|
||||
output_cost_per_token=output_rate,
|
||||
cache_read_input_token_cost=custom_cost_per_token.get("cache_read_input_token_cost", input_rate),
|
||||
cache_read_input_token_cost=cache_read_rate,
|
||||
cache_read_input_audio_token_cost=cache_read_rate,
|
||||
cache_creation_input_token_cost=cache_creation_rate,
|
||||
cache_creation_input_token_cost_above_1hr=cache_creation_rate,
|
||||
output_cost_per_reasoning_token=output_rate,
|
||||
|
|
@ -1413,6 +1454,11 @@ def _cost_map_billed_rates(
|
|||
completion_base_cost=completion_base_cost,
|
||||
current_time=billing_time,
|
||||
)
|
||||
audio_cache_read_rate: Final = _get_cost_per_unit(
|
||||
model_info,
|
||||
_get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier),
|
||||
None,
|
||||
)
|
||||
multiplier: Final = (
|
||||
_get_regional_uplift_multiplier(model_info, data_residency)
|
||||
* get_vertex_regional_endpoint_uplift(model_info, vertex_location)
|
||||
|
|
@ -1422,6 +1468,9 @@ def _cost_map_billed_rates(
|
|||
input_cost_per_token=prompt_base_cost,
|
||||
output_cost_per_token=completion_base_cost,
|
||||
cache_read_input_token_cost=cache_read_cost_rate,
|
||||
cache_read_input_audio_token_cost=(
|
||||
audio_cache_read_rate if audio_cache_read_rate is not None else cache_read_cost_rate
|
||||
),
|
||||
cache_creation_input_token_cost=cache_creation_cost_rate,
|
||||
cache_creation_input_token_cost_above_1hr=cache_creation_cost_above_1hr_rate,
|
||||
output_cost_per_reasoning_token=reasoning_rate,
|
||||
|
|
@ -1494,7 +1543,9 @@ def get_token_type_cost_breakdown(
|
|||
if rates is None:
|
||||
return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
|
||||
|
||||
cache_read_tokens, cache_creation_tokens, cache_creation_token_details = _cache_token_counts(usage)
|
||||
cache_read_tokens, cached_audio_tokens, cache_creation_tokens, cache_creation_token_details = _cache_token_counts(
|
||||
usage
|
||||
)
|
||||
cache_creation_cost: Final = (
|
||||
float(cache_creation_tokens) * rates.cache_creation_input_token_cost
|
||||
if custom_cost_per_token is not None
|
||||
|
|
@ -1507,7 +1558,10 @@ def get_token_type_cost_breakdown(
|
|||
)
|
||||
return TokenTypeCostBreakdown(
|
||||
reasoning_cost=float(_reasoning_token_count(usage)) * rates.output_cost_per_reasoning_token,
|
||||
cache_read_cost=float(cache_read_tokens) * rates.cache_read_input_token_cost,
|
||||
cache_read_cost=(
|
||||
float(cache_read_tokens - cached_audio_tokens) * rates.cache_read_input_token_cost
|
||||
+ float(cached_audio_tokens) * rates.cache_read_input_audio_token_cost
|
||||
),
|
||||
cache_creation_cost=cache_creation_cost,
|
||||
rates=rates,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from typing_extensions import ReadOnly
|
|||
|
||||
import litellm
|
||||
from litellm._logging import redact_internal_details_from_client_message, verbose_logger
|
||||
from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -35,9 +36,6 @@ else:
|
|||
CLIENT_CONNECTION_CLASS = Any
|
||||
|
||||
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY: Final = "realtime_session_success_logged"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BackendClose:
|
||||
code: int
|
||||
|
|
@ -1153,6 +1151,7 @@ class RealTimeStreaming:
|
|||
self._logging_worker.ensure_initialized_and_enqueue(
|
||||
self.logging_obj.dispatch_failure_handlers(error, traceback.format_exc(), prefer_async_handlers=True)
|
||||
)
|
||||
self.logging_obj.model_call_details[REALTIME_SESSION_FAILURE_LOGGED_KEY] = True
|
||||
|
||||
@staticmethod
|
||||
def _detect_beta_header(websocket: ScopedWebSocket) -> bool:
|
||||
|
|
|
|||
|
|
@ -7,13 +7,21 @@ This uses aws_sdk_bedrock_runtime for bidirectional streaming with Nova Sonic.
|
|||
import asyncio
|
||||
import contextlib
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from typing import Final, Protocol
|
||||
from collections.abc import AsyncIterator, Mapping, MutableMapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, NoReturn, Protocol
|
||||
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY,
|
||||
BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY,
|
||||
BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY,
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
)
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
|
@ -28,6 +36,32 @@ from .transformation import BedrockRealtimeConfig
|
|||
_CLIENT_MODALITIES_ADAPTER: Final[TypeAdapter["list[str] | None"]] = TypeAdapter(list[str] | None)
|
||||
_CLIENT_MESSAGE_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
|
||||
_EMPTY_JSON_OBJECT: Final[Mapping[str, JsonValue]] = MappingProxyType({})
|
||||
|
||||
_BEDROCK_STREAM_ERROR_STATUS: Final[Mapping[str, int]] = MappingProxyType(
|
||||
{
|
||||
"AccessDeniedException": 403,
|
||||
"ConflictException": 400,
|
||||
"InternalServerException": 500,
|
||||
"ModelErrorException": 424,
|
||||
"ModelNotReadyException": 429,
|
||||
"ModelStreamErrorException": 424,
|
||||
"ModelTimeoutException": 408,
|
||||
"ResourceNotFoundException": 404,
|
||||
"ServiceQuotaExceededException": 400,
|
||||
"ServiceUnavailableException": 503,
|
||||
"ThrottlingException": 429,
|
||||
"ValidationException": 400,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _as_bedrock_error(error: BaseException) -> BaseException:
|
||||
status_code: Final = _BEDROCK_STREAM_ERROR_STATUS.get(type(error).__name__)
|
||||
if status_code is None:
|
||||
return error
|
||||
return BedrockError(status_code=status_code, message=f"{type(error).__name__}: {error}")
|
||||
|
||||
|
||||
def _json_dict(value: JsonValue) -> dict[str, JsonValue]:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
|
@ -51,6 +85,8 @@ def _should_log_event(openai_message: Mapping[str, object]) -> bool:
|
|||
class RealtimeClientWebSocket(Protocol):
|
||||
"""The client-facing websocket surface the realtime bridge talks to."""
|
||||
|
||||
scope: MutableMapping[str, object] # mutable-ok: the ASGI scope is the per-connection state store
|
||||
|
||||
async def receive_text(self) -> str: ...
|
||||
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
|
@ -85,6 +121,81 @@ class BedrockBidirectionalStream(Protocol):
|
|||
async def await_output(self) -> tuple[object, BedrockOutputStream]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BridgeOutcome:
|
||||
logged_events: tuple[OpenAIRealtimeEvents, ...]
|
||||
provider_failure: BaseException | None
|
||||
client_disconnected: bool
|
||||
|
||||
|
||||
async def _client_messages(client_ws: RealtimeClientWebSocket, initial_message: str | None) -> AsyncIterator[str]:
|
||||
if initial_message is not None:
|
||||
yield initial_message
|
||||
while True:
|
||||
try:
|
||||
yield await client_ws.receive_text()
|
||||
except Exception as e: # noqa: BLE001 # any receive failure means the client is gone
|
||||
verbose_proxy_logger.debug("Client to Bedrock forwarding ended: %s", e, exc_info=True)
|
||||
return
|
||||
|
||||
|
||||
def _pending_session_update(scope: Mapping[str, object]) -> str | None:
|
||||
"""A fallback attempt on the same websocket replays the session.update the failed attempt never acked."""
|
||||
if scope.get(BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY) is True:
|
||||
committed_failure: Final = scope.get(BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY)
|
||||
raise BedrockError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"Bedrock realtime session already committed to a provider stream; it cannot be replayed"
|
||||
+ (f". The committed stream failed with: {committed_failure}" if committed_failure else "")
|
||||
),
|
||||
)
|
||||
pending: Final = scope.get(BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY)
|
||||
return pending if isinstance(pending, str) else None
|
||||
|
||||
|
||||
def _raise_provider_failure(scope: MutableMapping[str, object], failure: BaseException) -> NoReturn:
|
||||
error: Final = _as_bedrock_error(failure)
|
||||
verbose_proxy_logger.error("Bedrock Realtime: provider stream failed: %s", _redact_string(str(error)))
|
||||
if scope.get(BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY) is True:
|
||||
scope[BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY] = _redact_string(str(error))
|
||||
raise error from failure
|
||||
|
||||
|
||||
def _parse_client_message(message: str) -> Mapping[str, JsonValue]:
|
||||
try:
|
||||
return _json_dict(_CLIENT_MESSAGE_ADAPTER.validate_json(message))
|
||||
except ValueError:
|
||||
return _EMPTY_JSON_OBJECT
|
||||
|
||||
|
||||
async def _ack_session_update(
|
||||
client_ws: RealtimeClientWebSocket,
|
||||
bedrock_stream: BedrockBidirectionalStream,
|
||||
transformation_config: BedrockRealtimeConfig,
|
||||
model: str,
|
||||
logging_obj: LiteLLMLogging | None,
|
||||
parsed_client_message: Mapping[str, JsonValue],
|
||||
) -> bool:
|
||||
"""Ack the client's session.update once Bedrock accepted the stream; False means the client is gone."""
|
||||
await bedrock_stream.await_output()
|
||||
client_ws.scope.pop(BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY, None)
|
||||
client_ws.scope[BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY] = True # rebind-ok: scope outlives the attempt
|
||||
if logging_obj is None:
|
||||
return True
|
||||
requested_modalities: Final = _CLIENT_MODALITIES_ADAPTER.validate_python(
|
||||
_json_dict(parsed_client_message.get("session")).get("modalities")
|
||||
)
|
||||
try:
|
||||
await client_ws.send_text(
|
||||
json.dumps(transformation_config.session_updated_event(model, logging_obj, requested_modalities))
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # any send failure means the client is gone
|
||||
verbose_proxy_logger.debug("Client to Bedrock forwarding ended: %s", e, exc_info=True)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
class BedrockRealtime(BaseAWSLLM):
|
||||
"""Handler for Bedrock Nova Sonic realtime speech-to-speech API."""
|
||||
|
||||
|
|
@ -132,6 +243,8 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
except ImportError:
|
||||
raise ImportError("Missing aws_sdk_bedrock_runtime. Install with: pip install aws-sdk-bedrock-runtime")
|
||||
|
||||
pending_session_update: Final = _pending_session_update(websocket.scope)
|
||||
|
||||
# Get AWS region
|
||||
if aws_region_name is None:
|
||||
optional_params: Final = {
|
||||
|
|
@ -190,90 +303,106 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
|
||||
transformation_config: Final = BedrockRealtimeConfig()
|
||||
|
||||
try:
|
||||
# Initialize the bidirectional stream
|
||||
bedrock_stream: Final = await open_bidirectional_stream()
|
||||
bedrock_stream: Final = await open_bidirectional_stream()
|
||||
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established")
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established")
|
||||
|
||||
if pending_session_update is None:
|
||||
await websocket.send_text(json.dumps(transformation_config.session_created_event(model, logging_obj)))
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: sent session.created to client on connect")
|
||||
|
||||
# Track state for transformation
|
||||
session_state: Final[RealtimeResponseTransformInput] = {
|
||||
"current_output_item_id": None,
|
||||
"current_response_id": None,
|
||||
"current_conversation_id": None,
|
||||
"current_delta_chunks": None,
|
||||
"current_item_chunks": None,
|
||||
"current_delta_type": None,
|
||||
"session_configuration_request": None,
|
||||
}
|
||||
# Track state for transformation
|
||||
session_state: Final[RealtimeResponseTransformInput] = {
|
||||
"current_output_item_id": None,
|
||||
"current_response_id": None,
|
||||
"current_conversation_id": None,
|
||||
"current_delta_chunks": None,
|
||||
"current_item_chunks": None,
|
||||
"current_delta_type": None,
|
||||
"session_configuration_request": None,
|
||||
}
|
||||
|
||||
# Create tasks for bidirectional forwarding
|
||||
client_to_bedrock_task: Final = asyncio.create_task(
|
||||
self._forward_client_to_bedrock(
|
||||
websocket,
|
||||
bedrock_stream,
|
||||
transformation_config,
|
||||
model,
|
||||
session_state,
|
||||
logging_obj,
|
||||
outcome: Final = await self._bridge(
|
||||
websocket,
|
||||
bedrock_stream,
|
||||
transformation_config,
|
||||
model,
|
||||
session_state,
|
||||
logging_obj,
|
||||
initial_message=pending_session_update,
|
||||
)
|
||||
|
||||
logged_events: Final = (
|
||||
*outcome.logged_events,
|
||||
*(
|
||||
leftover_event
|
||||
for leftover_event in transformation_config.leftover_usage_done_events()
|
||||
if _should_log_event(leftover_event)
|
||||
),
|
||||
)
|
||||
if logged_events:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
logging_obj.dispatch_success_handlers(
|
||||
list(logged_events), # mutable-ok: realtime spend logging requires a list result
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
|
||||
|
||||
async def forward_bedrock_and_collect_logged_events() -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
return tuple(
|
||||
[
|
||||
event
|
||||
async for event in self._forward_bedrock_to_client(
|
||||
bedrock_stream,
|
||||
websocket,
|
||||
transformation_config,
|
||||
model,
|
||||
logging_obj,
|
||||
session_state,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
bedrock_to_client_task: Final = asyncio.create_task(forward_bedrock_and_collect_logged_events())
|
||||
|
||||
# Wait for both tasks to complete
|
||||
await asyncio.gather(
|
||||
client_to_bedrock_task,
|
||||
bedrock_to_client_task,
|
||||
return_exceptions=True,
|
||||
if outcome.provider_failure is None:
|
||||
return
|
||||
if outcome.client_disconnected:
|
||||
verbose_proxy_logger.debug(
|
||||
"Bedrock Realtime: stream failed after the client disconnected: %s", outcome.provider_failure
|
||||
)
|
||||
return
|
||||
_raise_provider_failure(websocket.scope, outcome.provider_failure)
|
||||
|
||||
forwarded_logged_events: Final = (
|
||||
bedrock_to_client_task.result()
|
||||
if not bedrock_to_client_task.cancelled() and bedrock_to_client_task.exception() is None
|
||||
else ()
|
||||
)
|
||||
logged_events: Final = (
|
||||
*forwarded_logged_events,
|
||||
*(
|
||||
leftover_event
|
||||
for leftover_event in transformation_config.leftover_usage_done_events()
|
||||
if _should_log_event(leftover_event)
|
||||
),
|
||||
)
|
||||
if logged_events:
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
logging_obj.dispatch_success_handlers(
|
||||
list(logged_events), # mutable-ok: realtime spend logging requires a list result
|
||||
prefer_async_handlers=True,
|
||||
)
|
||||
)
|
||||
async def _bridge(
|
||||
self,
|
||||
websocket: RealtimeClientWebSocket,
|
||||
bedrock_stream: BedrockBidirectionalStream,
|
||||
transformation_config: BedrockRealtimeConfig,
|
||||
model: str,
|
||||
session_state: RealtimeResponseTransformInput,
|
||||
logging_obj: LiteLLMLogging,
|
||||
initial_message: str | None,
|
||||
) -> _BridgeOutcome:
|
||||
"""Run both forwarding directions until the client leaves or either side fails."""
|
||||
logged: Final[list[OpenAIRealtimeEvents]] = [] # mutable-ok: events forwarded before a failure are still spend
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error in BedrockRealtime.async_realtime: %s", e)
|
||||
try:
|
||||
await websocket.close(code=1011, reason=_redact_string(f"Internal error: {e}"))
|
||||
except Exception:
|
||||
pass
|
||||
raise
|
||||
async def collect_logged_events() -> None:
|
||||
async for event in self._forward_bedrock_to_client(
|
||||
bedrock_stream, websocket, transformation_config, model, logging_obj, session_state
|
||||
):
|
||||
logged.append(event)
|
||||
|
||||
client_task: Final = asyncio.create_task(
|
||||
self._forward_client_to_bedrock(
|
||||
websocket, bedrock_stream, transformation_config, model, session_state, logging_obj, initial_message
|
||||
)
|
||||
)
|
||||
bedrock_task: Final = asyncio.create_task(collect_logged_events())
|
||||
|
||||
await asyncio.wait((client_task, bedrock_task), return_when=asyncio.FIRST_COMPLETED)
|
||||
client_disconnected: Final = (
|
||||
client_task.done() and not client_task.cancelled() and client_task.exception() is None
|
||||
)
|
||||
client_task.cancel()
|
||||
bedrock_task.cancel()
|
||||
client_outcome, bedrock_outcome = await asyncio.gather(client_task, bedrock_task, return_exceptions=True)
|
||||
|
||||
return _BridgeOutcome(
|
||||
logged_events=tuple(logged),
|
||||
provider_failure=(
|
||||
client_outcome
|
||||
if isinstance(client_outcome, Exception)
|
||||
else bedrock_outcome
|
||||
if isinstance(bedrock_outcome, Exception)
|
||||
else None
|
||||
),
|
||||
client_disconnected=client_disconnected,
|
||||
)
|
||||
|
||||
async def _forward_client_to_bedrock(
|
||||
self,
|
||||
|
|
@ -283,8 +412,12 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
model: str,
|
||||
session_state: RealtimeResponseTransformInput,
|
||||
logging_obj: LiteLLMLogging | None = None,
|
||||
):
|
||||
"""Forward messages from client WebSocket to Bedrock stream."""
|
||||
initial_message: str | None = None,
|
||||
) -> None:
|
||||
"""Forward messages from client WebSocket to Bedrock stream.
|
||||
|
||||
Returns once the client is gone; provider failures (input stream or readiness) propagate to the caller.
|
||||
"""
|
||||
from aws_sdk_bedrock_runtime.models import (
|
||||
BidirectionalInputPayloadPart,
|
||||
InvokeModelWithBidirectionalStreamInputChunk,
|
||||
|
|
@ -299,41 +432,28 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
verbose_proxy_logger.debug("Bedrock Realtime: Sent to Bedrock: %s", bedrock_message[:200])
|
||||
|
||||
try:
|
||||
while True:
|
||||
# Receive message from client
|
||||
message = await client_ws.receive_text()
|
||||
async for message in _client_messages(client_ws, initial_message):
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Received from client: %s", message[:200])
|
||||
parsed_client_message = _parse_client_message(message)
|
||||
is_session_update = _json_str(parsed_client_message.get("type")) == "session.update"
|
||||
if is_session_update:
|
||||
client_ws.scope[BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY] = (
|
||||
message # rebind-ok: scope outlives the attempt
|
||||
)
|
||||
|
||||
# Transform OpenAI format to Bedrock format
|
||||
transformed_messages = transformation_config.transform_realtime_request(
|
||||
message=message,
|
||||
model=model,
|
||||
session_configuration_request=session_state.get("session_configuration_request"),
|
||||
)
|
||||
|
||||
# Send transformed messages to Bedrock
|
||||
for bedrock_message in transformed_messages:
|
||||
await send_to_bedrock(bedrock_message)
|
||||
|
||||
if logging_obj is not None:
|
||||
client_message_type: str | None = None
|
||||
requested_modalities: list[str] | None = None
|
||||
with contextlib.suppress(Exception):
|
||||
parsed_client_message = _json_dict(_CLIENT_MESSAGE_ADAPTER.validate_json(message))
|
||||
client_message_type = _json_str(parsed_client_message.get("type"))
|
||||
if client_message_type == "session.update":
|
||||
requested_modalities = _CLIENT_MODALITIES_ADAPTER.validate_python(
|
||||
_json_dict(parsed_client_message.get("session")).get("modalities")
|
||||
)
|
||||
if client_message_type == "session.update":
|
||||
await client_ws.send_text(
|
||||
json.dumps(
|
||||
transformation_config.session_updated_event(model, logging_obj, requested_modalities)
|
||||
)
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Client to Bedrock forwarding ended: %s", e, exc_info=True)
|
||||
if is_session_update and not await _ack_session_update(
|
||||
client_ws, bedrock_stream, transformation_config, model, logging_obj, parsed_client_message
|
||||
):
|
||||
break
|
||||
finally:
|
||||
for close_message in transformation_config.session_close_messages():
|
||||
with contextlib.suppress(Exception):
|
||||
await send_to_bedrock(close_message)
|
||||
|
|
@ -349,68 +469,71 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
logging_obj: LiteLLMLogging,
|
||||
session_state: RealtimeResponseTransformInput,
|
||||
) -> AsyncIterator[OpenAIRealtimeEvents]:
|
||||
"""Forward messages from Bedrock to the client, yielding the ones to record for spend logging."""
|
||||
try:
|
||||
while True:
|
||||
# Receive from Bedrock
|
||||
output = await bedrock_stream.await_output()
|
||||
result = await output[1].receive()
|
||||
"""Forward messages from Bedrock to the client, yielding the ones to record for spend logging.
|
||||
|
||||
if result is None:
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Bedrock stream ended")
|
||||
break
|
||||
Provider failures propagate to the caller; the client websocket is only closed on a normal stream end.
|
||||
"""
|
||||
|
||||
payload_bytes = result.value.bytes_ if result.value else None
|
||||
if payload_bytes:
|
||||
bedrock_response = payload_bytes.decode("utf-8")
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Received from Bedrock: %s", bedrock_response[:200])
|
||||
|
||||
# Transform Bedrock format to OpenAI format
|
||||
realtime_response_transform_input: RealtimeResponseTransformInput = {
|
||||
"current_output_item_id": session_state.get("current_output_item_id"),
|
||||
"current_response_id": session_state.get("current_response_id"),
|
||||
"current_conversation_id": session_state.get("current_conversation_id"),
|
||||
"current_delta_chunks": session_state.get("current_delta_chunks"),
|
||||
"current_item_chunks": session_state.get("current_item_chunks"),
|
||||
"current_delta_type": session_state.get("current_delta_type"),
|
||||
"session_configuration_request": session_state.get("session_configuration_request"),
|
||||
}
|
||||
|
||||
transformed_response = transformation_config.transform_realtime_response(
|
||||
message=bedrock_response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
realtime_response_transform_input=realtime_response_transform_input,
|
||||
)
|
||||
|
||||
# Update session state
|
||||
session_state.update(
|
||||
{
|
||||
"current_output_item_id": transformed_response.get("current_output_item_id"),
|
||||
"current_response_id": transformed_response.get("current_response_id"),
|
||||
"current_conversation_id": transformed_response.get("current_conversation_id"),
|
||||
"current_delta_chunks": transformed_response.get("current_delta_chunks"),
|
||||
"current_item_chunks": transformed_response.get("current_item_chunks"),
|
||||
"current_delta_type": transformed_response.get("current_delta_type"),
|
||||
"session_configuration_request": transformed_response.get("session_configuration_request"),
|
||||
}
|
||||
)
|
||||
|
||||
# Send transformed messages to client
|
||||
response_value = transformed_response["response"]
|
||||
openai_messages = response_value if isinstance(response_value, list) else (response_value,)
|
||||
for openai_message in openai_messages:
|
||||
message_json = json.dumps(openai_message)
|
||||
await client_ws.send_text(message_json)
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Sent to client: %s", message_json[:200])
|
||||
if _should_log_event(openai_message):
|
||||
yield openai_message
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Bedrock to client forwarding ended: %s", e, exc_info=True)
|
||||
finally:
|
||||
# Close the client WebSocket
|
||||
async def send_to_client(message_json: str) -> bool:
|
||||
try:
|
||||
await client_ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
await client_ws.send_text(message_json)
|
||||
except Exception as e: # noqa: BLE001 # any send failure means the client is gone
|
||||
verbose_proxy_logger.debug("Bedrock to client forwarding ended: %s", e, exc_info=True)
|
||||
return False
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Sent to client: %s", message_json[:200])
|
||||
return True
|
||||
|
||||
output: Final = await bedrock_stream.await_output()
|
||||
while True:
|
||||
result = await output[1].receive()
|
||||
|
||||
if result is None:
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Bedrock stream ended")
|
||||
with contextlib.suppress(Exception):
|
||||
await client_ws.close()
|
||||
return
|
||||
|
||||
payload_bytes = result.value.bytes_ if result.value else None
|
||||
if payload_bytes:
|
||||
bedrock_response = payload_bytes.decode("utf-8")
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Received from Bedrock: %s", bedrock_response[:200])
|
||||
|
||||
# Transform Bedrock format to OpenAI format
|
||||
realtime_response_transform_input: RealtimeResponseTransformInput = {
|
||||
"current_output_item_id": session_state.get("current_output_item_id"),
|
||||
"current_response_id": session_state.get("current_response_id"),
|
||||
"current_conversation_id": session_state.get("current_conversation_id"),
|
||||
"current_delta_chunks": session_state.get("current_delta_chunks"),
|
||||
"current_item_chunks": session_state.get("current_item_chunks"),
|
||||
"current_delta_type": session_state.get("current_delta_type"),
|
||||
"session_configuration_request": session_state.get("session_configuration_request"),
|
||||
}
|
||||
|
||||
transformed_response = transformation_config.transform_realtime_response(
|
||||
message=bedrock_response,
|
||||
model=model,
|
||||
logging_obj=logging_obj,
|
||||
realtime_response_transform_input=realtime_response_transform_input,
|
||||
)
|
||||
|
||||
# Update session state
|
||||
session_state.update(
|
||||
{
|
||||
"current_output_item_id": transformed_response.get("current_output_item_id"),
|
||||
"current_response_id": transformed_response.get("current_response_id"),
|
||||
"current_conversation_id": transformed_response.get("current_conversation_id"),
|
||||
"current_delta_chunks": transformed_response.get("current_delta_chunks"),
|
||||
"current_item_chunks": transformed_response.get("current_item_chunks"),
|
||||
"current_delta_type": transformed_response.get("current_delta_type"),
|
||||
"session_configuration_request": transformed_response.get("session_configuration_request"),
|
||||
}
|
||||
)
|
||||
|
||||
# Send transformed messages to client
|
||||
response_value = transformed_response["response"]
|
||||
openai_messages = response_value if isinstance(response_value, list) else (response_value,)
|
||||
for openai_message in openai_messages:
|
||||
if not await send_to_client(json.dumps(openai_message)):
|
||||
return
|
||||
if _should_log_event(openai_message):
|
||||
yield openai_message
|
||||
|
|
|
|||
|
|
@ -5513,7 +5513,8 @@
|
|||
},
|
||||
"azure/gpt-realtime-2025-08-28": {
|
||||
"cache_creation_input_audio_token_cost": 4e-06,
|
||||
"cache_read_input_token_cost": 4e-06,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-03-02",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
|
|
@ -5546,7 +5547,8 @@
|
|||
},
|
||||
"azure/gpt-realtime-1.5-2026-02-23": {
|
||||
"cache_creation_input_audio_token_cost": 4e-06,
|
||||
"cache_read_input_token_cost": 4e-06,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-08-24",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
|
|
@ -5683,6 +5685,7 @@
|
|||
},
|
||||
"azure/gpt-realtime-mini": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-07,
|
||||
|
|
@ -5715,6 +5718,7 @@
|
|||
},
|
||||
"azure/gpt-realtime-mini-2025-10-06": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-07,
|
||||
|
|
@ -7409,6 +7413,80 @@
|
|||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/gpt-chat-latest": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"deprecation_date": "2026-12-02",
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"medium"
|
||||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": 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
|
||||
},
|
||||
"azure/chat-latest": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"deprecation_date": "2026-12-02",
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"medium"
|
||||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": 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
|
||||
},
|
||||
"azure/us/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
@ -7675,6 +7753,43 @@
|
|||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/us/gpt-chat-latest": {
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"deprecation_date": "2026-12-02",
|
||||
"input_cost_per_token": 5.5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"medium"
|
||||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": 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
|
||||
},
|
||||
"azure/eu/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
@ -23053,34 +23168,6 @@
|
|||
"output_cost_per_token": 0.0,
|
||||
"source": "https://fireworks.ai/pricing"
|
||||
},
|
||||
"friendliai/meta-llama-3.1-70b-instruct": {
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "friendliai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6e-07,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"friendliai/meta-llama-3.1-8b-instruct": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "friendliai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1e-07,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"friendliai/zai-org/GLM-5.3-Flash": {
|
||||
"litellm_provider": "friendliai",
|
||||
"max_input_tokens": 1048576,
|
||||
|
|
@ -32678,6 +32765,7 @@
|
|||
},
|
||||
"gpt-realtime": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-01-20",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
|
|
@ -32711,6 +32799,7 @@
|
|||
},
|
||||
"gpt-realtime-1.5": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
|
|
@ -32847,6 +32936,7 @@
|
|||
"gpt-realtime-mini": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"deprecation_date": "2027-01-20",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_token": 6e-07,
|
||||
|
|
@ -32878,6 +32968,7 @@
|
|||
},
|
||||
"gpt-realtime-2025-08-28": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-01-20",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
|
|
@ -57625,8 +57716,9 @@
|
|||
},
|
||||
{
|
||||
"name": "claude-adaptive-thinking",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"description": "Claude at version 4.6 or higher, in any id shape that contains claude-<family>-: minors 4.6 through 4.99, any later major-minor, and bare 5+ majors so a new family shaped like claude-fable-5 matches. Requiring the claude- prefix keeps non-Claude names such as team-sonnet-5-1 out. The minor is capped at two digits so an 8-digit date suffix such as claude-opus-4-20250514 is never read as a >= 4.6 minor. Turns on adaptive thinking for new versions and new families with no code change.",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|[5-9](?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"fill_missing_for_providers": ["anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai-anthropic_models"],
|
||||
"description": "Claude at version 4.6 or higher, in any id shape that contains claude-<family>-: minors 4.6 through 4.99, any later major-minor, and bare 5+ majors so a new family shaped like claude-fable-5 matches. Requiring the claude- prefix keeps non-Claude names such as team-sonnet-5-1 out. The minor is capped at two digits so an 8-digit date suffix such as claude-opus-4-20250514 is never read as a >= 4.6 minor. Two-digit majors are deliberately not matched so ids like claude-opus-41 (4.1) are not read as major 41. Turns on adaptive thinking for new versions and new families with no code change.",
|
||||
"model_info": {
|
||||
"supports_adaptive_thinking": true
|
||||
}
|
||||
|
|
@ -57634,6 +57726,7 @@
|
|||
{
|
||||
"name": "claude-legacy-thinking",
|
||||
"pattern": "claude-[a-z]+-4[-._]6(?!\\d)",
|
||||
"fill_missing_for_providers": ["anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai-anthropic_models"],
|
||||
"description": "Claude at version 4.6 exactly, in any id shape that contains claude-<family>-4-6 (dotted and underscored minors included, dated releases such as claude-sonnet-4-6-20260219 too). The 4.6 family is adaptive-thinking yet still accepts legacy thinking.type=enabled with budget_tokens, so the caller's hard budget cap is forwarded verbatim instead of being rewritten to an uncapped output_config.effort. The lookahead keeps two-digit minors such as 4-60 from matching. 4.7+ and 5+ majors reject the legacy shape and stay on the adaptive translation.",
|
||||
"model_info": {
|
||||
"supports_legacy_thinking": true
|
||||
|
|
@ -57649,8 +57742,9 @@
|
|||
},
|
||||
{
|
||||
"name": "claude-mid-conversation-system",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"description": "Claude at version 4.8 or higher, in any id shape that contains claude-<family>-: minors 4.8 through 4.99, any later major-minor, and bare 5+ majors so a new family like claude-fable-5 matches. Anthropic introduced mid-conversation system messages with Opus 4.8 and every newer Claude keeps them; 4.7 and below reject the system role inside messages.",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|[5-9](?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"fill_missing_for_providers": ["anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai-anthropic_models"],
|
||||
"description": "Claude at version 4.8 or higher, in any id shape that contains claude-<family>-: minors 4.8 through 4.99, any later major-minor, and bare 5+ majors so a new family like claude-fable-5 matches. Two-digit majors are deliberately not matched so ids like claude-opus-41 (4.1) are not read as major 41. Anthropic introduced mid-conversation system messages with Opus 4.8 and every newer Claude keeps them; 4.7 and below reject the system role inside messages.",
|
||||
"model_info": {
|
||||
"supports_mid_conversation_system": true
|
||||
}
|
||||
|
|
@ -57666,6 +57760,7 @@
|
|||
{
|
||||
"name": "openai-reasoning-family-baseline",
|
||||
"pattern": "^(?!.*search-api)(?:[a-z0-9_.-]+/)*(?:ft:)?(?:o[1-9]\\d*(?![a-z0-9])|gpt-[5-9](?:\\.\\d+)?(?![0-9.])|(?:gpt-\\d+(?:\\.\\d+)?(?:-[a-z0-9]+)*-)?(?:codex|deep-research|chat-latest)(?![a-z0-9]))",
|
||||
"fill_missing_for_providers": ["azure", "azure_ai", "openai"],
|
||||
"description": "OpenAI reasoning families by id shape, under any provider namespace and with an optional ft: prefix: the o-series (o1, o3-pro, o4-mini), gpt-5 through gpt-9 majors including dotted minors and suffixed variants (gpt-5.5-cyber, gpt-6-astra), and the codex, deep-research and chat-latest lines when standalone or on a gpt base. gpt-5-search-api is excluded because it is a search-only surface. Every model here is a reasoning model, and the Responses API drops the caller's reasoning param for any mapped OpenAI model whose info lacks supports_reasoning, so an id the registry has not named yet keeps its reasoning settings instead of silently losing them. Rules lose to exact entries. Carries no mode and no pricing, so cost stays on the standard unpriced behavior.",
|
||||
"model_info": {
|
||||
"supports_reasoning": true
|
||||
|
|
|
|||
|
|
@ -2637,9 +2637,13 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="max file size in MB for /v1/files uploads, for any purpose, if a file is larger than this size it will be rejected before being forwarded to the provider",
|
||||
)
|
||||
allowed_file_extensions: tuple[str, ...] | None = Field(
|
||||
None,
|
||||
description="the only file extensions (e.g. ['.jsonl', '.pdf', '.txt']) accepted on /v1/files uploads, for any purpose, matched case-insensitively against the uploaded filename. Files with any other extension, or none, are rejected. An empty list rejects every upload. Unset means no allowlist is applied",
|
||||
)
|
||||
blocked_file_extensions: tuple[str, ...] | None = Field(
|
||||
None,
|
||||
description="file extensions (e.g. ['.exe', '.sh']) rejected on /v1/files uploads, for any purpose, matched case-insensitively against the uploaded filename",
|
||||
description="file extensions (e.g. ['.exe', '.sh']) rejected on /v1/files uploads, for any purpose, matched case-insensitively against the uploaded filename. Deprecated in favour of allowed_file_extensions; still enforced, after the allowlist, when set",
|
||||
)
|
||||
max_response_size_mb: int | None = Field(
|
||||
None,
|
||||
|
|
@ -4839,6 +4843,14 @@ class JWTIssuerConfig(BaseModel):
|
|||
default=None,
|
||||
description="Issuer-specific claim path to normalize into LiteLLM's end-user id.",
|
||||
)
|
||||
virtual_key_claim_field: str | None = Field(
|
||||
default=None,
|
||||
description="Issuer-specific claim path used for the virtual key mapping lookup. Falls back to the global field.",
|
||||
)
|
||||
unregistered_jwt_client_behavior: UnregisteredJWTClientBehavior | None = Field(
|
||||
default=None,
|
||||
description="Issuer-specific policy when the virtual key claim has no mapping. Falls back to the global policy.",
|
||||
)
|
||||
|
||||
model_config = {
|
||||
"extra": "forbid",
|
||||
|
|
@ -5065,6 +5077,28 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def get_issuer_config(self, issuer: str | None) -> JWTIssuerConfig | None:
|
||||
if issuer is None or self.issuers is None:
|
||||
return None
|
||||
return next((config for config in self.issuers if config.issuer == issuer), None)
|
||||
|
||||
def is_virtual_key_mapping_configured(self) -> bool:
|
||||
if self.virtual_key_claim_field is not None:
|
||||
return True
|
||||
return any(config.virtual_key_claim_field is not None for config in self.issuers or ())
|
||||
|
||||
def get_virtual_key_claim_field(self, issuer: str | None) -> str | None:
|
||||
issuer_config: Final = self.get_issuer_config(issuer)
|
||||
if issuer_config is not None and issuer_config.virtual_key_claim_field is not None:
|
||||
return issuer_config.virtual_key_claim_field
|
||||
return self.virtual_key_claim_field
|
||||
|
||||
def get_unregistered_jwt_client_behavior(self, issuer: str | None) -> UnregisteredJWTClientBehavior:
|
||||
issuer_config: Final = self.get_issuer_config(issuer)
|
||||
if issuer_config is not None and issuer_config.unregistered_jwt_client_behavior is not None:
|
||||
return issuer_config.unregistered_jwt_client_behavior
|
||||
return self.unregistered_jwt_client_behavior
|
||||
|
||||
|
||||
class PrismaCompatibleUpdateDBModel(TypedDict, total=False):
|
||||
model_name: str
|
||||
|
|
|
|||
|
|
@ -123,6 +123,7 @@ from litellm.repositories.table_repositories import (
|
|||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.router import Router
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
||||
|
|
@ -327,9 +328,23 @@ _safe_json_loads_obj: Final = _typed_json_loads(safe_json_loads)
|
|||
last_db_access_time: Final = LimitedSizeOrderedDict(max_size=100)
|
||||
db_cache_expiry: Final = DEFAULT_IN_MEMORY_TTL # refresh every 5s
|
||||
|
||||
_TEAM_MEMBERSHIP_INFLIGHT_MAX: Final = 10000
|
||||
_team_membership_inflight: Final = LimitedSizeOrderedDict(max_size=_TEAM_MEMBERSHIP_INFLIGHT_MAX)
|
||||
|
||||
|
||||
class _TeamMembershipCacheMiss:
|
||||
__slots__ = ()
|
||||
|
||||
|
||||
_TEAM_MEMBERSHIP_CACHE_MISS: Final = _TeamMembershipCacheMiss()
|
||||
|
||||
all_routes: Final = LiteLLMRoutes.openai_routes.value + LiteLLMRoutes.management_routes.value
|
||||
|
||||
|
||||
def _membership_from_shared_load(result: object) -> LiteLLM_TeamMembership | None:
|
||||
return result if isinstance(result, LiteLLM_TeamMembership) else None
|
||||
|
||||
|
||||
def _log_budget_lookup_failure(entity: str, error: Exception) -> None:
|
||||
"""
|
||||
Log a warning when budget lookup fails; cache will not be populated.
|
||||
|
|
@ -887,6 +902,22 @@ async def common_checks(
|
|||
and (route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route))
|
||||
)
|
||||
|
||||
membership_user_id: Final = (
|
||||
valid_token.user_id if valid_token is not None and (bool(_model) or not skip_all_budget_checks) else None
|
||||
)
|
||||
team_membership_loaded: Final = team_object is not None and membership_user_id is not None
|
||||
loaded_team_membership: Final = (
|
||||
await get_team_membership(
|
||||
user_id=membership_user_id,
|
||||
team_id=team_object.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if team_object is not None and membership_user_id is not None
|
||||
else None
|
||||
)
|
||||
|
||||
unpriced_models: Final = (
|
||||
_unpriced_models_in_request(model=_model, llm_router=llm_router)
|
||||
if litellm.block_requests_for_models_without_pricing and RouteChecks.is_llm_api_route(route=route)
|
||||
|
|
@ -936,6 +967,8 @@ async def common_checks(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_membership=loaded_team_membership,
|
||||
team_membership_loaded=team_membership_loaded,
|
||||
)
|
||||
|
||||
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
|
||||
|
|
@ -987,6 +1020,8 @@ async def common_checks(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_membership=loaded_team_membership,
|
||||
team_membership_loaded=team_membership_loaded,
|
||||
)
|
||||
|
||||
# Run before apply_key_tags_pre_auth injects key metadata.tags into request_body.
|
||||
|
|
@ -1096,6 +1131,8 @@ async def common_checks(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_membership=loaded_team_membership,
|
||||
team_membership_loaded=team_membership_loaded,
|
||||
),
|
||||
_check_end_user_budget(end_user_obj=end_user_object, route=route)
|
||||
if end_user_object is not None and end_user_object.litellm_budget_table is not None
|
||||
|
|
@ -2141,7 +2178,76 @@ async def get_tag_object(
|
|||
return tag_objects.get(tag_name)
|
||||
|
||||
|
||||
def _membership_from_cached_payload(
|
||||
cached: object,
|
||||
) -> LiteLLM_TeamMembership | None | _TeamMembershipCacheMiss:
|
||||
if cached is None:
|
||||
return _TEAM_MEMBERSHIP_CACHE_MISS
|
||||
if cached == NO_TEAM_MEMBERSHIP_SENTINEL:
|
||||
return None
|
||||
cached_membership: Final = CacheCodec.deserialize(cached, model_type=LiteLLM_TeamMembership)
|
||||
return cached_membership if cached_membership is not None else _TEAM_MEMBERSHIP_CACHE_MISS
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def _fetch_team_membership_from_db(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Span | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> LiteLLM_TeamMembership | None:
|
||||
_ = parent_otel_span, proxy_logging_obj
|
||||
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique(
|
||||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
membership: Final = None if response is None else LiteLLM_TeamMembership.model_validate(response.dict())
|
||||
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
|
||||
if membership is None:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_key,
|
||||
value=NO_TEAM_MEMBERSHIP_SENTINEL,
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
)
|
||||
else:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_key,
|
||||
value=membership,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
)
|
||||
return membership
|
||||
|
||||
|
||||
async def _load_team_membership_on_cache_miss(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
cache_key: str,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
) -> LiteLLM_TeamMembership | None:
|
||||
try:
|
||||
redis_cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
redis_membership: Final = _membership_from_cached_payload(redis_cached)
|
||||
if not isinstance(redis_membership, _TeamMembershipCacheMiss):
|
||||
return redis_membership
|
||||
|
||||
return await _fetch_team_membership_from_db(
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("Error getting team membership")
|
||||
return None
|
||||
|
||||
|
||||
async def get_team_membership(
|
||||
user_id: str,
|
||||
team_id: str,
|
||||
|
|
@ -2155,54 +2261,42 @@ async def get_team_membership(
|
|||
|
||||
Do a isolated check for team membership vs. doing a combined key + team + user + team-membership check, as key might come in frequently for different users/teams. Larger call will slowdown query time. This way we get to cache the constant (key/team/user info) and only update based on the changing value (team membership).
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_TeamMembership
|
||||
|
||||
if prisma_client is None:
|
||||
raise Exception("No db connected")
|
||||
|
||||
if user_id is None or team_id is None:
|
||||
return None
|
||||
|
||||
_key: Final = team_membership_reservation_cache_key(user_id=user_id, team_id=team_id)
|
||||
|
||||
# check if in cache
|
||||
cached: Final[object] = await user_api_key_cache.async_get_cache(key=_key)
|
||||
if cached == NO_TEAM_MEMBERSHIP_SENTINEL:
|
||||
return None
|
||||
cached_membership_obj: Final = CacheCodec.deserialize(cached, model_type=LiteLLM_TeamMembership)
|
||||
if cached_membership_obj is not None:
|
||||
return cached_membership_obj
|
||||
l1_cached: Final[object] = await user_api_key_cache.async_get_cache(key=_key, local_only=True)
|
||||
l1_membership: Final = _membership_from_cached_payload(l1_cached)
|
||||
if not isinstance(l1_membership, _TeamMembershipCacheMiss):
|
||||
return l1_membership
|
||||
|
||||
# else, check db
|
||||
try:
|
||||
response: Final = await _dictable_table(TeamMembershipRepository(prisma_client)).find_unique(
|
||||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}},
|
||||
include={"litellm_budget_table": True},
|
||||
inflight: Final[object] = _team_membership_inflight.get(_key)
|
||||
if isinstance(inflight, asyncio.Task):
|
||||
return _membership_from_shared_load(await asyncio.shield(inflight))
|
||||
|
||||
if prisma_client is None:
|
||||
raise Exception("No db connected")
|
||||
|
||||
task: Final = asyncio.ensure_future(
|
||||
_load_team_membership_on_cache_miss(
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
cache_key=_key,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
_team_membership_inflight[_key] = task
|
||||
|
||||
if response is None:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_key,
|
||||
value=NO_TEAM_MEMBERSHIP_SENTINEL,
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
)
|
||||
return None
|
||||
def _clear_inflight(_done: object) -> None:
|
||||
if _team_membership_inflight.get(_key) is task:
|
||||
_team_membership_inflight.pop(_key, None)
|
||||
|
||||
_response: Final = LiteLLM_TeamMembership.model_validate(response.dict())
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_key,
|
||||
value=_response,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
)
|
||||
|
||||
return _response
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"Error getting team membership for user_id: %s, team_id: %s",
|
||||
user_id,
|
||||
team_id,
|
||||
)
|
||||
return None
|
||||
task.add_done_callback(_clear_inflight)
|
||||
return _membership_from_shared_load(await asyncio.shield(task))
|
||||
|
||||
|
||||
def model_in_access_group(model: str, team_models: list[str] | None, llm_router: Router | None) -> bool:
|
||||
|
|
@ -2375,13 +2469,6 @@ async def _backfill_null_user_email(
|
|||
return updated_row
|
||||
|
||||
|
||||
class UserNotFoundError(ValueError):
|
||||
"""The user row is provably absent, as opposed to merely unreadable, so a caller that reads a missing row as no user-level limits can key on it without also swallowing a database that would not answer."""
|
||||
|
||||
def __init__(self, user_id: str) -> None:
|
||||
super().__init__(f"User doesn't exist in db. 'user_id'={user_id}. Create user via `/user/new` call.")
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def get_user_object(
|
||||
user_id: str | None,
|
||||
|
|
@ -2668,6 +2755,12 @@ async def invalidate_team_member_spend_state(
|
|||
publish_auth_cache_invalidation,
|
||||
)
|
||||
|
||||
inflight: Final[object] = _team_membership_inflight.pop(
|
||||
team_membership_reservation_cache_key(user_id=user_id, team_id=team_id), None
|
||||
)
|
||||
if isinstance(inflight, asyncio.Task) and inflight is not asyncio.current_task():
|
||||
await asyncio.wait((inflight,))
|
||||
|
||||
if new_spend is not None:
|
||||
from litellm.proxy.proxy_server import SPEND_DB_FLOOR_CACHE_TTL_SECONDS, spend_counter_cache
|
||||
|
||||
|
|
@ -4122,18 +4215,21 @@ async def _team_member_granted_models(
|
|||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_membership: LiteLLM_TeamMembership | None = None,
|
||||
team_membership_loaded: bool = False,
|
||||
) -> Sequence[str]:
|
||||
"""The member's own ``allowed_models`` scope; empty when the member is not narrowed below the team."""
|
||||
if team_object is None or valid_token.user_id is None:
|
||||
return ()
|
||||
|
||||
team_membership: Final = await get_team_membership(
|
||||
user_id=valid_token.user_id,
|
||||
team_id=team_object.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if not team_membership_loaded:
|
||||
team_membership = await get_team_membership(
|
||||
user_id=valid_token.user_id,
|
||||
team_id=team_object.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
return () if team_membership is None else _member_allowed_models(team_membership)
|
||||
|
||||
|
||||
|
|
@ -4169,6 +4265,8 @@ async def _granted_model_lists(
|
|||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_membership: LiteLLM_TeamMembership | None = None,
|
||||
team_membership_loaded: bool = False,
|
||||
) -> tuple[Sequence[str], ...]:
|
||||
"""One model allowlist per level that participates in authorizing the request."""
|
||||
return (
|
||||
|
|
@ -4180,6 +4278,8 @@ async def _granted_model_lists(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_membership=team_membership,
|
||||
team_membership_loaded=team_membership_loaded,
|
||||
),
|
||||
project_object.models if project_object is not None else (),
|
||||
await _org_granted_models(
|
||||
|
|
@ -4274,6 +4374,8 @@ async def collect_matched_model_access_groups(
|
|||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_membership: LiteLLM_TeamMembership | None = None,
|
||||
team_membership_loaded: bool = False,
|
||||
) -> tuple[str, ...]:
|
||||
"""
|
||||
The budgeted model access groups that authorized this request, sorted and deduplicated.
|
||||
|
|
@ -4319,6 +4421,8 @@ async def collect_matched_model_access_groups(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_membership=team_membership,
|
||||
team_membership_loaded=team_membership_loaded,
|
||||
)
|
||||
for granted_model in granted_models
|
||||
)
|
||||
|
|
@ -4334,6 +4438,8 @@ async def stamp_matched_model_access_groups(
|
|||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_membership: LiteLLM_TeamMembership | None = None,
|
||||
team_membership_loaded: bool = False,
|
||||
) -> tuple[str, ...]:
|
||||
"""Record the groups that authorized this request on its auth object, for the post-call spend
|
||||
writer and the reservation counters, and hand them back for the budget check."""
|
||||
|
|
@ -4350,6 +4456,8 @@ async def stamp_matched_model_access_groups(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_membership=team_membership,
|
||||
team_membership_loaded=team_membership_loaded,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # fail-safe: attribution is spend telemetry, it must never break auth
|
||||
verbose_proxy_logger.debug("model access group attribution failed: %s", e)
|
||||
|
|
@ -5152,6 +5260,8 @@ async def _check_team_member_budget(
|
|||
prisma_client: PrismaClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_membership: LiteLLM_TeamMembership | None = None,
|
||||
team_membership_loaded: bool = False,
|
||||
):
|
||||
"""Check if team member is over their max budget within the team."""
|
||||
if (
|
||||
|
|
@ -5160,23 +5270,25 @@ async def _check_team_member_budget(
|
|||
and valid_token is not None
|
||||
and valid_token.user_id is not None
|
||||
):
|
||||
team_membership: Final = await get_team_membership(
|
||||
user_id=valid_token.user_id,
|
||||
team_id=team_object.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if not team_membership_loaded:
|
||||
team_membership = await get_team_membership(
|
||||
user_id=valid_token.user_id,
|
||||
team_id=team_object.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
loaded_membership = team_membership
|
||||
|
||||
# Per-member override wins; otherwise fall back to the team-level
|
||||
# default configured via team.metadata["team_member_budget_id"].
|
||||
team_member_budget: float | None = None
|
||||
if (
|
||||
team_membership is not None
|
||||
and team_membership.litellm_budget_table is not None
|
||||
and team_membership.litellm_budget_table.max_budget is not None
|
||||
loaded_membership is not None
|
||||
and loaded_membership.litellm_budget_table is not None
|
||||
and loaded_membership.litellm_budget_table.max_budget is not None
|
||||
):
|
||||
team_member_budget = team_membership.litellm_budget_table.max_budget
|
||||
team_member_budget = loaded_membership.litellm_budget_table.max_budget
|
||||
else:
|
||||
default_budget_id: Final = (team_object.metadata or {}).get("team_member_budget_id")
|
||||
if isinstance(default_budget_id, str):
|
||||
|
|
@ -5195,7 +5307,7 @@ async def _check_team_member_budget(
|
|||
team_member_budget = default_budget.max_budget
|
||||
|
||||
if team_member_budget is not None:
|
||||
team_member_spend = (team_membership.spend if team_membership is not None else 0.0) or 0.0
|
||||
team_member_spend = (loaded_membership.spend if loaded_membership is not None else 0.0) or 0.0
|
||||
|
||||
# Read from cross-pod counter (Redis-first) if available
|
||||
from litellm.proxy.proxy_server import get_current_spend
|
||||
|
|
@ -5224,6 +5336,8 @@ async def _check_team_member_model_access(
|
|||
prisma_client: Optional["PrismaClient"],
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_membership: LiteLLM_TeamMembership | None = None,
|
||||
team_membership_loaded: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
Check if a team member's per-member model scope allows access to the requested model.
|
||||
|
|
@ -5234,22 +5348,24 @@ async def _check_team_member_model_access(
|
|||
if valid_token.user_id is None or team_object.team_id is None:
|
||||
return
|
||||
|
||||
team_membership: Final = await get_team_membership(
|
||||
user_id=valid_token.user_id,
|
||||
team_id=team_object.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if not team_membership_loaded:
|
||||
team_membership = await get_team_membership(
|
||||
user_id=valid_token.user_id,
|
||||
team_id=team_object.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
loaded_membership = team_membership
|
||||
|
||||
if (
|
||||
team_membership is None
|
||||
or team_membership.litellm_budget_table is None
|
||||
or not team_membership.litellm_budget_table.allowed_models
|
||||
loaded_membership is None
|
||||
or loaded_membership.litellm_budget_table is None
|
||||
or not loaded_membership.litellm_budget_table.allowed_models
|
||||
):
|
||||
return # no per-member restriction — inherit team-level check
|
||||
|
||||
member_allowed_models: Final[list[str]] = team_membership.litellm_budget_table.allowed_models
|
||||
member_allowed_models: Final[list[str]] = loaded_membership.litellm_budget_table.allowed_models
|
||||
try:
|
||||
_can_object_call_model(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -28,11 +28,11 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
TeamNotFoundError,
|
||||
UserNotFoundError,
|
||||
get_team_membership,
|
||||
get_team_object,
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import Span
|
||||
|
|
|
|||
|
|
@ -138,6 +138,9 @@ class RouteChecks:
|
|||
# For llm_api_routes, also check registered pass-through endpoints
|
||||
################################################
|
||||
if allowed_route == "llm_api_routes":
|
||||
if route == "/auto_router/session" and RouteChecks._get_request_method(request) == "GET":
|
||||
return True
|
||||
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -996,9 +996,12 @@ async def _resolve_jwt_to_virtual_key(
|
|||
- Raises HTTPException: REJECT policy hit, missing claim under
|
||||
REJECT/AUTO_REGISTER, or other policy violations.
|
||||
"""
|
||||
virtual_key_claim_field: Final = jwt_handler.litellm_jwtauth.virtual_key_claim_field
|
||||
raw_issuer: Final = jwt_claims.get(JWTHandler.LITELLM_JWT_ISSUER_CLAIM)
|
||||
normalized_issuer: Final = raw_issuer if isinstance(raw_issuer, str) else None
|
||||
virtual_key_claim_field: Final = jwt_handler.litellm_jwtauth.get_virtual_key_claim_field(normalized_issuer)
|
||||
if virtual_key_claim_field is None:
|
||||
return None
|
||||
behavior: Final = jwt_handler.litellm_jwtauth.get_unregistered_jwt_client_behavior(normalized_issuer)
|
||||
|
||||
claim_value: Final = get_nested_value(
|
||||
data=jwt_claims,
|
||||
|
|
@ -1015,7 +1018,6 @@ async def _resolve_jwt_to_virtual_key(
|
|||
# simply by presenting a JWT that omits the configured field. For
|
||||
# AUTO_REGISTER there is no stable identity to map without a claim
|
||||
# value, so we deny rather than create a sentinel-keyed record.
|
||||
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
|
||||
if behavior in (
|
||||
UnregisteredJWTClientBehavior.REJECT,
|
||||
UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
|
|
@ -1030,7 +1032,13 @@ async def _resolve_jwt_to_virtual_key(
|
|||
return None
|
||||
|
||||
cache_key: Final = jwt_key_mapping_cache_key(virtual_key_claim_field, str(claim_value))
|
||||
cached_mapping: Final = await user_api_key_cache.async_get_cache(cache_key)
|
||||
raw_cached_mapping: Final = await user_api_key_cache.async_get_cache(cache_key)
|
||||
sentinel_written_by_this_policy: Final = behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER
|
||||
cached_mapping: Final = (
|
||||
None
|
||||
if raw_cached_mapping == _JWT_PROXY_ADMIN_SENTINEL and not sentinel_written_by_this_policy
|
||||
else raw_cached_mapping
|
||||
)
|
||||
|
||||
if cached_mapping == _JWT_PROXY_ADMIN_SENTINEL:
|
||||
# Previously resolved to a proxy admin via auth_builder; skip the
|
||||
|
|
@ -1039,7 +1047,6 @@ async def _resolve_jwt_to_virtual_key(
|
|||
return None
|
||||
|
||||
if cached_mapping == "__NO_MAPPING__":
|
||||
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
|
||||
if behavior == UnregisteredJWTClientBehavior.REJECT:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
@ -1102,8 +1109,6 @@ async def _resolve_jwt_to_virtual_key(
|
|||
)
|
||||
|
||||
# No mapping found (DB miss or no DB) — apply no-match policy.
|
||||
behavior = jwt_handler.litellm_jwtauth.unregistered_jwt_client_behavior
|
||||
|
||||
if behavior == UnregisteredJWTClientBehavior.REJECT:
|
||||
# Cache the miss before raising so repeated rejections are served from
|
||||
# cache and don't re-query the DB on every request.
|
||||
|
|
@ -1483,7 +1488,7 @@ async def _user_api_key_auth_builder(
|
|||
# unnecessary DB queries in auth_builder
|
||||
do_standard_jwt_auth = True
|
||||
pending_auto_register: _PendingAutoRegister | None = None
|
||||
if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None:
|
||||
if jwt_handler.litellm_jwtauth.is_virtual_key_mapping_configured():
|
||||
# Decode JWT to get claims without running full auth_builder
|
||||
jwt_claims: dict | None
|
||||
if jwt_handler.litellm_jwtauth.oidc_userinfo_enabled and not is_jwt:
|
||||
|
|
|
|||
|
|
@ -585,7 +585,9 @@ LiteLLM ████████░░░░░░░░░░░░░░
|
|||
Claude Opus 5 ████████████████████████ $0.38
|
||||
```
|
||||
|
||||
The routed model comes from Claude Code's own transcript, so it only names the tier model when the auto-router deployment sets `return_raw_model_name: true` (the `lite autoroute` wizard does); otherwise it shows the alias you requested. The cost lines come from `GET /auto_router/session?session_id=...`, which any virtual key may call for its own sessions, and are cached for five seconds under a per-user `$TMPDIR/litellm-statusline-<uid>` directory. The baseline is the priciest model in the router's hardest tier, the same counterfactual the auto-router's savings reports use. `lite unconfigure claude` removes the `statusLine` entry only while it still points at that script.
|
||||
After the first response, the status line uses the latest routed model recorded by `GET /auto_router/session?session_id=...`, so it can show the tier model even when the transcript contains the router alias. If no session record is available, it falls back to Claude Code's transcript. Session records and costs are cached for five seconds under a per-user `$TMPDIR/litellm-statusline-<uid>` directory. The gateway records turns asynchronously, so the display can briefly lag a completed turn. Any virtual key may read its own sessions. The baseline is the priciest model in the router's hardest tier, the same counterfactual the auto-router's savings reports use. `lite unconfigure claude` removes the `statusLine` entry only while it still points at that script
|
||||
|
||||
After upgrading the CLI, rerun your original `lite configure claude` command with the same gateway, key and model choice to refresh `~/.litellm/statusline.py`. Keep any explicit `--model` value: omitting it removes the earlier model pin. Package upgrades alone do not refresh this installed copy
|
||||
|
||||
`lite codex` registers the same script as a Codex `Stop` hook for the launch, so after each turn Codex prints the same block as a system message. Codex asks once to trust the hook; the answer is remembered for later launches.
|
||||
|
||||
|
|
|
|||
|
|
@ -7,19 +7,17 @@ status refresh (about every 300ms while typing), so the proxy is asked at most o
|
|||
TTL per session and every other refresh is served from a small on-disk cache that holds
|
||||
only the proxy's answer, never the key.
|
||||
|
||||
Claude Code pipes a JSON payload on stdin (session_id, transcript_path, model); the routed
|
||||
model is the `message.model` of the latest foreground assistant line in the transcript,
|
||||
which is the proxy's response `model` field. That only names the tier model when the
|
||||
auto-router deployment sets `return_raw_model_name: true`; otherwise it is the alias the
|
||||
client requested. Codex pipes its Stop event instead (hook_event_name, session_id) and has
|
||||
no transcript to read, so the routed model comes from the proxy's session record and the
|
||||
result is printed as a `systemMessage` for the transcript. The proxy key is read from the
|
||||
agent's own environment (the static token `lite configure claude` writes); nothing here
|
||||
spawns a credential helper.
|
||||
Claude Code pipes a JSON payload on stdin (session_id, transcript_path, model). After the
|
||||
first foreground assistant response, the routed model comes from the proxy's session
|
||||
record, falling back to the latest foreground assistant `message.model` in the transcript
|
||||
when no record is available. Codex pipes its Stop event instead (hook_event_name, session_id)
|
||||
and prints the session record as a `systemMessage` for the transcript. The proxy key is read
|
||||
from the agent's own environment (the static token `lite configure claude` writes); nothing
|
||||
here spawns a credential helper.
|
||||
|
||||
Cost figures come from GET /auto_router/session on the proxy, which reads the per-session
|
||||
rollup written by the spend flush. That flush is asynchronous, so a turn's cost lands a
|
||||
second or two after the turn; the cache TTL absorbs it.
|
||||
The routed model and cost figures come from GET /auto_router/session on the proxy, which
|
||||
reads the per-session rollup written by the asynchronous spend flush. The record and cache
|
||||
can briefly lag a completed turn.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -348,7 +346,8 @@ def status_line(
|
|||
if not session_id or not credentials.usable:
|
||||
return render(label, None, config_dir, color_enabled(env))
|
||||
session: Final = load_session(credentials, session_id, cache_dir, fetch)
|
||||
return render(label, session, config_dir, color_enabled(env))
|
||||
routed_label: Final = model_label(session.last_model, config_dir) if session is not None else label
|
||||
return render(routed_label, session, config_dir, color_enabled(env))
|
||||
|
||||
|
||||
def codex_stop_message(
|
||||
|
|
|
|||
|
|
@ -3452,15 +3452,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# a failed request reports no timing, matching /v1/chat/completions
|
||||
read_timing_from_logging_obj=False,
|
||||
)
|
||||
# Extract headers from exception - check both e.headers and e.response.headers
|
||||
headers = getattr(e, "headers", None) or {}
|
||||
if not headers:
|
||||
# Try to get headers from e.response.headers (httpx.Response)
|
||||
_response: Final = attribute_of(e, "response")
|
||||
if _response is not None:
|
||||
_response_headers: Final = getattr(_response, "headers", None)
|
||||
if _response_headers:
|
||||
headers = get_response_headers(dict(_response_headers))
|
||||
_response_headers: Final = getattr(_response, "headers", None) if _response is not None else None
|
||||
_provider_headers: Final = _response_headers or getattr(e, "litellm_response_headers", None)
|
||||
if _provider_headers:
|
||||
headers = get_response_headers(dict(_provider_headers))
|
||||
headers.update(custom_headers)
|
||||
|
||||
# Call response headers hook for failure
|
||||
|
|
|
|||
|
|
@ -0,0 +1,10 @@
|
|||
import os
|
||||
from collections.abc import Mapping
|
||||
|
||||
|
||||
def should_hide_default_credentials_hint(general_settings: Mapping[str, object]) -> bool:
|
||||
return (
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
or bool(os.getenv("UI_PASSWORD"))
|
||||
)
|
||||
|
|
@ -4,6 +4,7 @@ from typing import Final
|
|||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint
|
||||
from litellm.types.proxy.discovery_endpoints.ui_discovery_endpoints import (
|
||||
UiDiscoveryEndpoints,
|
||||
)
|
||||
|
|
@ -23,10 +24,7 @@ async def get_ui_config():
|
|||
or general_settings.get("auto_redirect_ui_login_to_sso", False) is True
|
||||
)
|
||||
admin_ui_disabled: Final = os.getenv("DISABLE_ADMIN_UI", "false").lower() == "true"
|
||||
hide_default_credentials_hint: Final = bool(
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
)
|
||||
hide_default_credentials_hint: Final = should_hide_default_credentials_hint(general_settings)
|
||||
|
||||
sso_configured: Final = has_user_setup_sso()
|
||||
|
||||
|
|
|
|||
|
|
@ -6,11 +6,10 @@ from typing import Final
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import PROXY_LLM_PROVIDER_FALLBACK
|
||||
from litellm.types.router import ModelGroupInfo
|
||||
from litellm.types.utils import PriorityReservationDict
|
||||
|
||||
PROXY_LLM_PROVIDER_FALLBACK: Final = "litellm_proxy"
|
||||
|
||||
|
||||
def resolve_llm_provider_for_rate_limit(
|
||||
model: str | None,
|
||||
|
|
|
|||
|
|
@ -156,6 +156,7 @@ from litellm.repositories.verification_token_repository import (
|
|||
VerificationTokenRepository,
|
||||
)
|
||||
from litellm.router import Router
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
from litellm.types.proxy.management_endpoints.common_daily_activity import (
|
||||
SpendAnalyticsPaginatedResponse,
|
||||
)
|
||||
|
|
@ -181,6 +182,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
if TYPE_CHECKING:
|
||||
from prisma import Prisma
|
||||
from prisma import models as prisma_models
|
||||
from prisma import types as prisma_types
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
|
@ -4889,6 +4891,26 @@ async def _get_org_admin_org_ids(
|
|||
return org_ids if org_ids else None
|
||||
|
||||
|
||||
async def _get_user_team_ids_from_db(
|
||||
user_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> tuple[str, ...]:
|
||||
try:
|
||||
user: Final = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=True,
|
||||
)
|
||||
except UserNotFoundError:
|
||||
return ()
|
||||
return tuple(user.teams or ()) if user is not None else ()
|
||||
|
||||
|
||||
async def _build_team_list_where_conditions(
|
||||
prisma_client: PrismaClient,
|
||||
team_id: str | None,
|
||||
|
|
@ -4899,12 +4921,16 @@ async def _build_team_list_where_conditions(
|
|||
search: str | None = None,
|
||||
search_team_id_match: TeamIdSearchMatch = "exact",
|
||||
org_admin_org_ids: list[str] | None = None,
|
||||
own_team_ids: tuple[str, ...] = (),
|
||||
user_api_key_cache: UserApiKeyCache | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Build where conditions for team list query.
|
||||
|
||||
An org admin listing their own teams sees the union of the teams in the
|
||||
orgs they administer and `own_team_ids`, the teams they are a member of.
|
||||
|
||||
Returns None when the query is guaranteed to yield no results (e.g. user
|
||||
has no team memberships), allowing the caller to skip the DB round-trip.
|
||||
"""
|
||||
|
|
@ -4927,6 +4953,11 @@ async def _build_team_list_where_conditions(
|
|||
|
||||
if organization_id:
|
||||
where_conditions["organization_id"] = organization_id
|
||||
elif org_admin_org_ids is not None and own_team_ids:
|
||||
org_or_membership_scope: Final[prisma_types.LiteLLM_TeamTableWhereInput] = {
|
||||
"OR": [{"organization_id": {"in": org_admin_org_ids}}, {"team_id": {"in": list(own_team_ids)}}]
|
||||
}
|
||||
where_conditions["AND"] = [org_or_membership_scope]
|
||||
elif org_admin_org_ids is not None:
|
||||
# Org admin: always scope to their orgs, even when filtering by user_id.
|
||||
where_conditions["organization_id"] = {"in": org_admin_org_ids}
|
||||
|
|
@ -5058,66 +5089,72 @@ async def _enforce_list_team_v2_access(
|
|||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> tuple[str | None, list[str] | None]:
|
||||
) -> tuple[str | None, list[str] | None, tuple[str, ...]]:
|
||||
"""Enforce access control for list_team_v2.
|
||||
|
||||
- Proxy admins and admin viewers can query any teams.
|
||||
- Org admins can query teams within their organizations.
|
||||
- Org admins can query teams within their organizations, plus the teams
|
||||
they are a member of when listing their own teams.
|
||||
- Regular users can only query their own teams.
|
||||
|
||||
Returns the (possibly overridden) user_id and org_admin_org_ids.
|
||||
Returns the (possibly overridden) user_id, org_admin_org_ids and, for an
|
||||
org admin's own query, the caller's own team ids.
|
||||
"""
|
||||
is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
org_admin_org_ids: list[str] | None = None
|
||||
caller_user_id: Final = user_api_key_dict.user_id
|
||||
|
||||
if is_proxy_admin:
|
||||
return user_id, org_admin_org_ids
|
||||
return user_id, None, ()
|
||||
|
||||
# Always check org admin status so that even own-queries see
|
||||
# the full set of organisation teams, not just direct memberships.
|
||||
if user_api_key_dict.user_id:
|
||||
org_admin_org_ids = await _get_org_admin_org_ids(
|
||||
user_id=user_api_key_dict.user_id,
|
||||
org_admin_org_ids: Final = (
|
||||
await _get_org_admin_org_ids(
|
||||
user_id=caller_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if caller_user_id
|
||||
else None
|
||||
)
|
||||
|
||||
if org_admin_org_ids is not None:
|
||||
if caller_user_id and org_admin_org_ids is not None:
|
||||
# Org admin: validate org_id filter if provided
|
||||
if organization_id and organization_id not in org_admin_org_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "You can only view teams within your organizations."},
|
||||
)
|
||||
# When the caller is an org admin querying their own teams (or no
|
||||
# specific user), null out user_id so that
|
||||
# _build_team_list_where_conditions scopes only by organization_id
|
||||
# — org admins should see all teams in their orgs, not just teams
|
||||
# they are a direct member of. Keep user_id when the org admin
|
||||
# explicitly queries a *different* user's teams.
|
||||
if user_id is None or user_id == user_api_key_dict.user_id:
|
||||
user_id = None
|
||||
is_own_query: Final = user_id is None or user_id == caller_user_id
|
||||
own_team_ids: Final = (
|
||||
await _get_user_team_ids_from_db(
|
||||
user_id=caller_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if is_own_query
|
||||
else ()
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"list_team_v2: org admin access for user=%s, org_ids=%s, user_id_filter=%s",
|
||||
user_api_key_dict.user_id,
|
||||
_sanitize_for_log(caller_user_id),
|
||||
org_admin_org_ids,
|
||||
user_id,
|
||||
_sanitize_for_log(None if is_own_query else user_id),
|
||||
)
|
||||
else:
|
||||
# Not an org admin — fall back to standard route check
|
||||
if not allowed_route_check_inside_route(user_api_key_dict=user_api_key_dict, requested_user_id=user_id):
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": f"Only admin users can query all teams/other teams. Your user role={user_api_key_dict.user_role}"
|
||||
},
|
||||
)
|
||||
# Regular user — auto-inject caller's user_id
|
||||
if user_id is None:
|
||||
user_id = user_api_key_dict.user_id
|
||||
return None if is_own_query else user_id, org_admin_org_ids, own_team_ids
|
||||
|
||||
return user_id, org_admin_org_ids
|
||||
# Not an org admin — fall back to standard route check
|
||||
if not allowed_route_check_inside_route(user_api_key_dict=user_api_key_dict, requested_user_id=user_id):
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": f"Only admin users can query all teams/other teams. Your user role={user_api_key_dict.user_role}"
|
||||
},
|
||||
)
|
||||
# Regular user — auto-inject caller's user_id
|
||||
return user_id if user_id is not None else caller_user_id, None, ()
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -5195,7 +5232,7 @@ async def list_team_v2(
|
|||
)
|
||||
|
||||
# --- Access control ---
|
||||
user_id, org_admin_org_ids = await _enforce_list_team_v2_access(
|
||||
user_id, org_admin_org_ids, own_team_ids = await _enforce_list_team_v2_access(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
user_id=user_id,
|
||||
organization_id=organization_id,
|
||||
|
|
@ -5227,6 +5264,7 @@ async def list_team_v2(
|
|||
search=search,
|
||||
search_team_id_match=search_team_id_match,
|
||||
org_admin_org_ids=org_admin_org_ids,
|
||||
own_team_ids=own_team_ids,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -5323,17 +5361,16 @@ async def _authorize_and_filter_teams(
|
|||
|
||||
- Proxy admins: all teams (or filtered by user_id if provided).
|
||||
- Org admins: teams from their orgs (scoped to user_id if provided).
|
||||
- Own query (user_id matches caller): teams the user is a member of.
|
||||
- Own query (user_id matches caller): teams the user is a member of, across all orgs.
|
||||
- Others: 401.
|
||||
"""
|
||||
is_proxy_admin: Final = _user_has_admin_view(user_api_key_dict)
|
||||
is_own_query: Final = (
|
||||
user_id is not None and user_api_key_dict.user_id is not None and user_api_key_dict.user_id == user_id
|
||||
)
|
||||
allowed_org_ids: list[str] | None = None
|
||||
|
||||
if not is_proxy_admin:
|
||||
is_own_query: Final = (
|
||||
user_id is not None and user_api_key_dict.user_id is not None and user_api_key_dict.user_id == user_id
|
||||
)
|
||||
|
||||
# Check if user is an org admin (even for own queries, so they see org teams)
|
||||
if user_api_key_dict.user_id is not None:
|
||||
caller_user: Final = await get_user_object(
|
||||
|
|
@ -5360,33 +5397,30 @@ async def _authorize_and_filter_teams(
|
|||
},
|
||||
)
|
||||
|
||||
if allowed_org_ids is not None:
|
||||
# Org admin: query DB for teams in their orgs
|
||||
if allowed_org_ids is not None and not is_own_query:
|
||||
org_teams: Final = await _raw_team_db(TeamRepository(prisma_client)).find_many(
|
||||
where={"organization_id": {"in": allowed_org_ids}},
|
||||
include={"litellm_model_table": True},
|
||||
)
|
||||
if not user_id:
|
||||
return list(org_teams)
|
||||
# Filter org teams to only those where the target user is a member
|
||||
return [
|
||||
team
|
||||
for team in org_teams
|
||||
if team.members_with_roles and any(m.get("user_id") == user_id for m in team.members_with_roles)
|
||||
]
|
||||
elif user_id:
|
||||
# Regular user: fetch all and filter by membership (Prisma can't filter JSON arrays)
|
||||
response: Final = await _raw_team_db(TeamRepository(prisma_client)).find_many(
|
||||
include={"litellm_model_table": True}
|
||||
)
|
||||
return [
|
||||
team
|
||||
for team in response
|
||||
if team.members_with_roles and any(m.get("user_id") == user_id for m in team.members_with_roles)
|
||||
]
|
||||
else:
|
||||
|
||||
response: Final = await _raw_team_db(TeamRepository(prisma_client)).find_many(include={"litellm_model_table": True})
|
||||
if not user_id:
|
||||
# Proxy admin: all teams
|
||||
return list(await _raw_team_db(TeamRepository(prisma_client)).find_many(include={"litellm_model_table": True}))
|
||||
return list(response)
|
||||
|
||||
# Prisma can't filter JSON arrays, so membership is filtered in Python
|
||||
return [
|
||||
team
|
||||
for team in response
|
||||
if team.members_with_roles and any(m.get("user_id") == user_id for m in team.members_with_roles)
|
||||
]
|
||||
|
||||
|
||||
@router.get("/team/list", tags=["team management"], dependencies=[Depends(user_api_key_auth)])
|
||||
|
|
|
|||
|
|
@ -101,6 +101,7 @@ from litellm.proxy.common_utils.admin_ui_utils import (
|
|||
admin_ui_disabled,
|
||||
show_missing_vars_in_env,
|
||||
)
|
||||
from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint
|
||||
from litellm.proxy.common_utils.html_forms.jwt_display_template import (
|
||||
jwt_display_template,
|
||||
)
|
||||
|
|
@ -1110,10 +1111,7 @@ async def google_login(
|
|||
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
hide_default_credentials_hint: Final = (
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
)
|
||||
hide_default_credentials_hint: Final = should_hide_default_credentials_hint(general_settings)
|
||||
form_response: Final = HTMLResponse(
|
||||
content=build_ui_login_form(
|
||||
show_deprecation_banner=True,
|
||||
|
|
|
|||
|
|
@ -78,6 +78,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
)
|
||||
from litellm.proxy.openai_files_endpoints.general_upload_validation import (
|
||||
MB,
|
||||
check_allowed_extension,
|
||||
check_blocked_extension,
|
||||
check_unsafe_filename,
|
||||
check_upload_file_size,
|
||||
|
|
@ -473,6 +474,11 @@ async def create_file(
|
|||
if general_size_failure is not None:
|
||||
raise_upload_validation_failure(general_size_failure)
|
||||
|
||||
allowed_extensions: Final = coerce_optional_str_list_setting(general_settings.get("allowed_file_extensions"))
|
||||
allowed_extension_failure: Final = check_allowed_extension(file.filename, allowed_extensions)
|
||||
if allowed_extension_failure is not None:
|
||||
raise_upload_validation_failure(allowed_extension_failure)
|
||||
|
||||
blocked_extensions: Final = coerce_optional_str_list_setting(general_settings.get("blocked_file_extensions"))
|
||||
blocked_extension_failure: Final = check_blocked_extension(file.filename, blocked_extensions)
|
||||
if blocked_extension_failure is not None:
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@
|
|||
Upload validation applied to every purpose at POST /v1/files.
|
||||
|
||||
batch_file_validation.py checks the JSONL shape of purpose="batch" uploads; this
|
||||
module applies the same fast-fail-before-forwarding shape (size cap, blocked
|
||||
extensions, path-traversal filenames) regardless of purpose.
|
||||
module applies the same fast-fail-before-forwarding shape (size cap, allowed and
|
||||
blocked extensions, path-traversal filenames) regardless of purpose.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -31,10 +31,9 @@ def coerce_optional_int_setting(raw: object) -> int | None:
|
|||
raise TypeError(f"expected an integer, got {raw!r}")
|
||||
|
||||
|
||||
def coerce_optional_str_list_setting(raw: object) -> tuple[str, ...]:
|
||||
"""A general_settings value declared as an optional list of strings, e.g. blocked_file_extensions."""
|
||||
def coerce_optional_str_list_setting(raw: object) -> tuple[str, ...] | None:
|
||||
if raw is None:
|
||||
return ()
|
||||
return None
|
||||
if not isinstance(raw, list) or not all(isinstance(item, str) for item in raw):
|
||||
raise TypeError(f"expected a list of strings, got {raw!r}")
|
||||
return tuple(raw)
|
||||
|
|
@ -46,6 +45,11 @@ class UploadedFileTooLarge:
|
|||
limit_mb: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class UploadedFileExtensionNotAllowed:
|
||||
extension: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class UploadedFileBlockedExtension:
|
||||
extension: str
|
||||
|
|
@ -56,7 +60,9 @@ class UploadedFileUnsafeFilename:
|
|||
filename: str
|
||||
|
||||
|
||||
UploadValidationFailure = UploadedFileTooLarge | UploadedFileBlockedExtension | UploadedFileUnsafeFilename
|
||||
UploadValidationFailure = (
|
||||
UploadedFileTooLarge | UploadedFileExtensionNotAllowed | UploadedFileBlockedExtension | UploadedFileUnsafeFilename
|
||||
)
|
||||
|
||||
|
||||
def _file_size_bytes(file_source: bytes | BinaryIO) -> int:
|
||||
|
|
@ -81,19 +87,35 @@ def check_upload_file_size(
|
|||
return None
|
||||
|
||||
|
||||
def _normalized_extension(filename: str | None) -> str:
|
||||
if not filename:
|
||||
return ""
|
||||
try:
|
||||
return Path(safe_filename(filename)).suffix.lower()
|
||||
except ValueError:
|
||||
return ""
|
||||
|
||||
|
||||
def check_allowed_extension(
|
||||
filename: str | None,
|
||||
allowed_extensions: tuple[str, ...] | None,
|
||||
) -> UploadedFileExtensionNotAllowed | None:
|
||||
if allowed_extensions is None:
|
||||
return None
|
||||
extension: Final = _normalized_extension(filename)
|
||||
normalized_allowed: Final = frozenset(item.lower() for item in allowed_extensions)
|
||||
if extension and extension in normalized_allowed:
|
||||
return None
|
||||
return UploadedFileExtensionNotAllowed(extension=extension)
|
||||
|
||||
|
||||
def check_blocked_extension(
|
||||
filename: str | None,
|
||||
blocked_extensions: tuple[str, ...],
|
||||
blocked_extensions: tuple[str, ...] | None,
|
||||
) -> UploadedFileBlockedExtension | None:
|
||||
if not blocked_extensions or not filename:
|
||||
if not blocked_extensions:
|
||||
return None
|
||||
try:
|
||||
extension: Final = Path(safe_filename(filename)).suffix.lower()
|
||||
except ValueError:
|
||||
return None
|
||||
# The uploaded name's extension is normalized above; blocked_extensions comes
|
||||
# straight from config.yaml or the DB and is normalized here too, so a
|
||||
# differently-cased entry (".EXE") still catches a lowercase upload.
|
||||
extension: Final = _normalized_extension(filename)
|
||||
normalized_blocked: Final = frozenset(item.lower() for item in blocked_extensions)
|
||||
if extension and extension in normalized_blocked:
|
||||
return UploadedFileBlockedExtension(extension=extension)
|
||||
|
|
@ -128,6 +150,17 @@ def raise_upload_validation_failure(failure: UploadValidationFailure) -> NoRetur
|
|||
param="file",
|
||||
code=413,
|
||||
)
|
||||
case UploadedFileExtensionNotAllowed(extension=extension):
|
||||
raise ProxyException(
|
||||
message=(
|
||||
(f"File extension '{extension}'" if extension else "A file without an extension")
|
||||
+ " is not in this proxy's allowed_file_extensions setting. "
|
||||
"The file was not forwarded to the provider."
|
||||
),
|
||||
type="invalid_request_error",
|
||||
param="file",
|
||||
code=400,
|
||||
)
|
||||
case UploadedFileBlockedExtension(extension=extension):
|
||||
raise ProxyException(
|
||||
message=(
|
||||
|
|
|
|||
|
|
@ -250,7 +250,7 @@ import litellm._redis
|
|||
from litellm import Router
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger, verbose_router_logger
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError
|
||||
from litellm.caching.redis_cache import RedisCircuitBreakerOpenError, is_redis_timeout_failure
|
||||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||||
from litellm.constants import (
|
||||
_REALTIME_BODY_CACHE_SIZE,
|
||||
|
|
@ -274,6 +274,8 @@ from litellm.constants import (
|
|||
PROXY_BUDGET_RESCHEDULER_MAX_TIME,
|
||||
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
|
||||
PROXY_CONFIG_RELOAD_INTERVAL_SECONDS,
|
||||
REALTIME_SESSION_FAILURE_LOGGED_KEY,
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG,
|
||||
USER_SPEND_ALERTS_JOB_ID,
|
||||
WEEKLY_SPEND_REPORT_JOB_ID,
|
||||
|
|
@ -365,6 +367,7 @@ from litellm.proxy.common_utils.healthy_model_filter import (
|
|||
get_hidden_unhealthy_model_names,
|
||||
is_healthy_only_listing_default,
|
||||
)
|
||||
from litellm.proxy.common_utils.html_forms.default_credentials_hint import should_hide_default_credentials_hint
|
||||
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
|
|
@ -3450,8 +3453,10 @@ async def _invalidate_spend_counter(counter_key: str):
|
|||
async def _apply_spend_counter_increments(pending: Sequence[PendingSpendIncrement]) -> None:
|
||||
try:
|
||||
await increment_spend_counters_pipeline(pending=pending)
|
||||
except RedisCircuitBreakerOpenError:
|
||||
return
|
||||
except Exception as e:
|
||||
if isinstance(e, RedisCircuitBreakerOpenError) or is_redis_timeout_failure(e):
|
||||
return
|
||||
raise
|
||||
|
||||
|
||||
async def increment_spend_counters_pipeline(pending: Sequence[PendingSpendIncrement]) -> None:
|
||||
|
|
@ -7065,6 +7070,9 @@ class ProxyConfig:
|
|||
if "max_file_size_mb" not in self._yaml_general_settings_keys:
|
||||
general_settings["max_file_size_mb"] = _general_settings.get("max_file_size_mb")
|
||||
|
||||
if "allowed_file_extensions" not in self._yaml_general_settings_keys:
|
||||
general_settings["allowed_file_extensions"] = _general_settings.get("allowed_file_extensions")
|
||||
|
||||
if "blocked_file_extensions" not in self._yaml_general_settings_keys:
|
||||
general_settings["blocked_file_extensions"] = _general_settings.get("blocked_file_extensions")
|
||||
|
||||
|
|
@ -11893,6 +11901,13 @@ async def _release_realtime_budget_reservation(user_api_key_dict: UserAPIKeyAuth
|
|||
)
|
||||
|
||||
|
||||
async def _release_realtime_max_parallel_slot(user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
release_like_http_disconnect: Final = (
|
||||
proxy_logging_obj._arelease_max_parallel_requests_on_disconnect # pyright: ignore[reportPrivateUsage] # shared
|
||||
)
|
||||
await release_like_http_disconnect(user_api_key_dict)
|
||||
|
||||
|
||||
async def _reject_realtime_session(
|
||||
websocket: WebSocket,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -11912,6 +11927,7 @@ async def _reject_realtime_session(
|
|||
await websocket.close(code=code, reason=reason)
|
||||
finally:
|
||||
await _release_realtime_budget_reservation(user_api_key_dict)
|
||||
await _release_realtime_max_parallel_slot(user_api_key_dict)
|
||||
|
||||
|
||||
@app.websocket("/openai/v1/realtime")
|
||||
|
|
@ -12015,6 +12031,9 @@ async def realtime_websocket_endpoint(
|
|||
websocket, user_api_key_dict, code=1011, reason="Pre-call error", error_message=str(e)
|
||||
)
|
||||
return
|
||||
except BaseException:
|
||||
await _release_realtime_max_parallel_slot(user_api_key_dict)
|
||||
raise
|
||||
|
||||
# Phase 2: route to upstream LLM.
|
||||
try:
|
||||
|
|
@ -12044,12 +12063,10 @@ async def realtime_websocket_endpoint(
|
|||
except Exception: # noqa: BLE001 # the lower layer may have closed the socket already; closing twice is not an error
|
||||
verbose_proxy_logger.debug("Could not close realtime client websocket; it is already gone")
|
||||
finally:
|
||||
from litellm.litellm_core_utils.realtime_streaming import (
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
)
|
||||
|
||||
if not litellm_logging_obj.model_call_details.get(REALTIME_SESSION_SUCCESS_LOGGED_KEY):
|
||||
await _release_realtime_budget_reservation(user_api_key_dict)
|
||||
if not litellm_logging_obj.model_call_details.get(REALTIME_SESSION_FAILURE_LOGGED_KEY):
|
||||
await _release_realtime_max_parallel_slot(user_api_key_dict)
|
||||
|
||||
|
||||
######################################################################
|
||||
|
|
@ -15033,6 +15050,13 @@ def _get_proxy_model_info(model: dict) -> dict:
|
|||
return _translate_model_name_for_response(model)
|
||||
|
||||
|
||||
def _model_info_json_response(data: Sequence[Mapping[str, object]] | Mapping[str, object]) -> Response:
|
||||
return Response(
|
||||
content=orjson.dumps({"data": data}, default=jsonable_encoder, option=orjson.OPT_NON_STR_KEYS),
|
||||
media_type="application/json",
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/model/info",
|
||||
tags=["model management"],
|
||||
|
|
@ -15080,7 +15104,7 @@ async def model_info_v1(
|
|||
`model_info.direct_access` when the proxy database is connected.
|
||||
|
||||
Returns:
|
||||
Returns a dictionary containing information about each model.
|
||||
A JSON response whose `data` list holds one entry per model.
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
|
|
@ -15128,7 +15152,7 @@ async def model_info_v1(
|
|||
deployment_dict=_deployment_info_dict,
|
||||
excluded_keys={"litellm_credential_name"},
|
||||
)
|
||||
return {"data": _deployment_info_dict}
|
||||
return _model_info_json_response(_deployment_info_dict)
|
||||
|
||||
if llm_model_list is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -15179,7 +15203,7 @@ async def model_info_v1(
|
|||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return {"data": single_model_list}
|
||||
return _model_info_json_response(single_model_list)
|
||||
|
||||
# Return router deployments (same source as /v2/model/info), not wildcard-
|
||||
# expanded model names from get_complete_model_list(). Team-scoped rows
|
||||
|
|
@ -15247,7 +15271,7 @@ async def model_info_v1(
|
|||
visible_models: Final = [model for model in all_models if model.get("model_name") not in hidden_names]
|
||||
|
||||
verbose_proxy_logger.debug("all_models: %s", visible_models)
|
||||
return {"data": visible_models}
|
||||
return _model_info_json_response(visible_models)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -15816,10 +15840,7 @@ async def fallback_login(request: Request):
|
|||
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
hide_default_credentials_hint: Final = (
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
)
|
||||
hide_default_credentials_hint: Final = should_hide_default_credentials_hint(general_settings)
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(
|
||||
show_deprecation_banner=False,
|
||||
|
|
@ -17036,6 +17057,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
|
|||
"max_request_size_mb": "Integer",
|
||||
"max_batch_file_size_mb": "Integer",
|
||||
"max_file_size_mb": "Integer",
|
||||
"allowed_file_extensions": "List",
|
||||
"blocked_file_extensions": "List",
|
||||
"max_response_size_mb": "Integer",
|
||||
"proxy_config_reload_interval_seconds": "Integer",
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
|||
from litellm.responses.litellm_completion_transformation.session_handler import (
|
||||
ResponsesSessionHandler,
|
||||
)
|
||||
from litellm.types.llms.base import CachedTokensDetails
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionAssistantMessage,
|
||||
|
|
@ -2816,27 +2817,24 @@ class LiteLLMCompletionResponsesConfig:
|
|||
# Translate prompt_tokens_details to input_tokens_details
|
||||
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details is not None:
|
||||
prompt_details: Final = usage.prompt_tokens_details
|
||||
input_details_dict: Final[dict[str, int]] = {}
|
||||
|
||||
if hasattr(prompt_details, "cached_tokens") and prompt_details.cached_tokens is not None:
|
||||
input_details_dict["cached_tokens"] = prompt_details.cached_tokens
|
||||
else:
|
||||
input_details_dict["cached_tokens"] = 0
|
||||
|
||||
if hasattr(prompt_details, "text_tokens") and prompt_details.text_tokens is not None:
|
||||
input_details_dict["text_tokens"] = prompt_details.text_tokens
|
||||
|
||||
if hasattr(prompt_details, "audio_tokens") and prompt_details.audio_tokens is not None:
|
||||
input_details_dict["audio_tokens"] = prompt_details.audio_tokens
|
||||
|
||||
cache_write_tokens = getattr(prompt_details, "cache_write_tokens", None) or getattr(
|
||||
cached_tokens_details: Final = getattr(prompt_details, "cached_tokens_details", None)
|
||||
cache_write_tokens: Final = getattr(prompt_details, "cache_write_tokens", None) or getattr(
|
||||
prompt_details, "cache_creation_tokens", None
|
||||
)
|
||||
if cache_write_tokens is not None:
|
||||
input_details_dict["cache_write_tokens"] = cache_write_tokens
|
||||
|
||||
if input_details_dict:
|
||||
response_usage.input_tokens_details = InputTokensDetails(**input_details_dict)
|
||||
cache_write_extra: Final[Mapping[str, int]] = (
|
||||
MappingProxyType({"cache_write_tokens": cache_write_tokens})
|
||||
if cache_write_tokens is not None
|
||||
else MappingProxyType({})
|
||||
)
|
||||
response_usage.input_tokens_details = InputTokensDetails(
|
||||
cached_tokens=prompt_details.cached_tokens if prompt_details.cached_tokens is not None else 0,
|
||||
text_tokens=prompt_details.text_tokens,
|
||||
audio_tokens=prompt_details.audio_tokens,
|
||||
cached_tokens_details=(
|
||||
cached_tokens_details if isinstance(cached_tokens_details, CachedTokensDetails) else None
|
||||
),
|
||||
**cache_write_extra,
|
||||
)
|
||||
|
||||
# Translate completion_tokens_details to output_tokens_details
|
||||
if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details is not None:
|
||||
|
|
|
|||
|
|
@ -1179,6 +1179,9 @@ class ResponseAPILoggingUtils:
|
|||
audio_tokens=getattr(response_api_usage.input_tokens_details, "audio_tokens", None),
|
||||
text_tokens=getattr(response_api_usage.input_tokens_details, "text_tokens", None),
|
||||
image_tokens=getattr(response_api_usage.input_tokens_details, "image_tokens", None),
|
||||
cached_tokens_details=getattr(
|
||||
response_api_usage.input_tokens_details, "cached_tokens_details", None
|
||||
),
|
||||
cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
|
||||
)
|
||||
completion_tokens_details: CompletionTokensDetailsWrapper | None = None
|
||||
|
|
|
|||
|
|
@ -96,7 +96,6 @@ from litellm.litellm_core_utils.request_timeout_resolver import (
|
|||
from litellm.litellm_core_utils.secret_redaction import redact_string
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import (
|
||||
SensitiveDataMasker,
|
||||
mask_credentials_in_payload,
|
||||
mask_sensitive_structure,
|
||||
)
|
||||
from litellm.litellm_core_utils.token_counter import offload_token_count
|
||||
|
|
@ -623,20 +622,6 @@ def _replay_live_router_model_cost() -> None:
|
|||
set_live_deployment_replay(_replay_live_router_model_cost)
|
||||
|
||||
|
||||
# Kwargs that carry no signal about the failed attempt, so log_retry drops them from a
|
||||
# breadcrumb entirely: the request payload, the proxy's snapshot of the inbound request (its body
|
||||
# aliases the live request metadata, earlier breadcrumbs included, so copying it would nest every
|
||||
# breadcrumb inside the next one), and the router-internal walk state. Credentials are handled
|
||||
# separately by mask_credentials_in_payload, which scrubs credential-named values from whatever
|
||||
# kwargs remain rather than trying to enumerate every credential-bearing key here.
|
||||
RETRY_BREADCRUMB_EXCLUDED_KWARGS: Final = frozenset(
|
||||
(
|
||||
"messages",
|
||||
"original_function",
|
||||
"attempted_targets",
|
||||
"proxy_server_request",
|
||||
)
|
||||
)
|
||||
RETRY_BREADCRUMB_LIMIT: Final = 4
|
||||
|
||||
|
||||
|
|
@ -1553,6 +1538,18 @@ class Router:
|
|||
return False
|
||||
return sum(len(self.model_name_to_deployment_indices.get(member) or ()) for member in group.models) > 1
|
||||
|
||||
def team_model_has_alternatives(self, deployment_id: str) -> bool:
|
||||
deployment: Final = self.get_deployment(model_id=deployment_id)
|
||||
if deployment is None:
|
||||
return False
|
||||
team_id: Final = deployment.model_info.team_id
|
||||
public_model_name: Final = deployment.model_info.team_public_model_name
|
||||
if team_id is None or public_model_name is None:
|
||||
return False
|
||||
sibling_indices: Final = self.team_model_to_deployment_indices.get((team_id, public_model_name)) or ()
|
||||
routable_siblings: Final = self._filter_blocked_deployments([self.model_list[idx] for idx in sibling_indices])
|
||||
return len(routable_siblings) > 1
|
||||
|
||||
_OVERRIDABLE_ROUTING_STRATEGIES: frozenset[str] = frozenset({"simple-shuffle", *_DEFAULT_SELECTOR_ATTR_BY_STRATEGY})
|
||||
|
||||
def _get_request_routing_strategy_override(self, request_kwargs: dict | None) -> str | None:
|
||||
|
|
@ -8374,31 +8371,30 @@ class Router:
|
|||
|
||||
def log_retry(self, kwargs: dict, e: Exception) -> dict:
|
||||
"""
|
||||
When a retry or fallback happens, log the details of the just failed model call - similar to Sentry breadcrumbing
|
||||
When a retry or fallback happens, record which model group, deployment and attempt just failed and why
|
||||
"""
|
||||
from litellm.types.router import RetryAttemptRecord
|
||||
|
||||
_metadata_var: Final = "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
|
||||
request_metadata: Final[Mapping[str, object]] = kwargs[_metadata_var]
|
||||
attempt_kwargs: Final = MappingProxyType(
|
||||
{k: v for k, v in kwargs.items() if k != _metadata_var and k not in RETRY_BREADCRUMB_EXCLUDED_KWARGS}
|
||||
)
|
||||
attempt_metadata: Final = MappingProxyType(
|
||||
{k: v for k, v in request_metadata.items() if k != "previous_models"}
|
||||
)
|
||||
previous_model: Final = MappingProxyType(
|
||||
{
|
||||
"exception_type": type(e).__name__,
|
||||
"exception_string": str(e),
|
||||
**attempt_kwargs,
|
||||
_metadata_var: attempt_metadata,
|
||||
}
|
||||
)
|
||||
model_group: Final = kwargs.get("model")
|
||||
model_info: Final = request_metadata.get("model_info")
|
||||
deployment_id: Final = model_info.get("id") if isinstance(model_info, Mapping) else None
|
||||
attempted_retries: Final = request_metadata.get("attempted_retries")
|
||||
attempt_record: Final[RetryAttemptRecord] = {
|
||||
"model_group": model_group if isinstance(model_group, str) else None,
|
||||
"deployment_id": deployment_id if isinstance(deployment_id, str) else None,
|
||||
"exception_type": type(e).__name__,
|
||||
"exception_string": str(e),
|
||||
"attempted_retries": attempted_retries if type(attempted_retries) is int else None,
|
||||
}
|
||||
earlier_breadcrumbs: Final = request_metadata.get("previous_models")
|
||||
kept_breadcrumbs: Final[tuple[object, ...]] = (
|
||||
tuple(earlier_breadcrumbs)[-(RETRY_BREADCRUMB_LIMIT - 1) :]
|
||||
if isinstance(earlier_breadcrumbs, (list, tuple))
|
||||
else ()
|
||||
)
|
||||
breadcrumbs: Final = (*kept_breadcrumbs, mask_credentials_in_payload(previous_model))
|
||||
breadcrumbs: Final = (*kept_breadcrumbs, attempt_record)
|
||||
kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict
|
||||
return kwargs
|
||||
|
||||
|
|
@ -13878,6 +13874,7 @@ class Router:
|
|||
cooldown_time=_cooldown_time,
|
||||
enable_pre_call_checks=self.enable_pre_call_checks,
|
||||
cooldown_list=_cooldown_list,
|
||||
model_ids=model_ids,
|
||||
)
|
||||
|
||||
if strategy == "simple-shuffle":
|
||||
|
|
@ -13910,6 +13907,7 @@ class Router:
|
|||
cooldown_time=_cooldown_time,
|
||||
enable_pre_call_checks=self.enable_pre_call_checks,
|
||||
cooldown_list=_cooldown_list,
|
||||
model_ids=model_ids,
|
||||
)
|
||||
self._override_selector_pre_call_check(strategy, strategy_selector, deployment)
|
||||
verbose_router_logger.info(
|
||||
|
|
@ -13987,6 +13985,11 @@ class Router:
|
|||
model=model,
|
||||
llm_provider="",
|
||||
)
|
||||
pass_through_model_ids: Final = tuple(
|
||||
deployment["model_info"]["id"]
|
||||
for deployment in pass_through_deployments
|
||||
if "id" in deployment.get("model_info", {})
|
||||
)
|
||||
|
||||
# 4. Apply health-check and cooldown filtering
|
||||
parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs)
|
||||
|
|
@ -14024,6 +14027,7 @@ class Router:
|
|||
cooldown_time=_cooldown_time,
|
||||
enable_pre_call_checks=self.enable_pre_call_checks,
|
||||
cooldown_list=_cooldown_list,
|
||||
model_ids=pass_through_model_ids,
|
||||
)
|
||||
|
||||
# 6. Apply load balancing strategy
|
||||
|
|
@ -14057,6 +14061,7 @@ class Router:
|
|||
cooldown_time=_cooldown_time,
|
||||
enable_pre_call_checks=self.enable_pre_call_checks,
|
||||
cooldown_list=_cooldown_list,
|
||||
model_ids=model_ids,
|
||||
)
|
||||
self._override_selector_pre_call_check(strategy, strategy_selector, deployment)
|
||||
|
||||
|
|
|
|||
|
|
@ -343,8 +343,9 @@ def _should_cooldown_deployment(
|
|||
model_group: Final = litellm_router_instance.get_model_group(id=deployment)
|
||||
is_single_deployment_model_group = False
|
||||
if model_group is not None and len(model_group) == 1:
|
||||
is_single_deployment_model_group = not litellm_router_instance.routing_group_has_alternatives(
|
||||
requested_model_group
|
||||
is_single_deployment_model_group = not (
|
||||
litellm_router_instance.routing_group_has_alternatives(requested_model_group)
|
||||
or litellm_router_instance.team_model_has_alternatives(deployment)
|
||||
)
|
||||
|
||||
## CHECK DEPLOYMENT-LEVEL POLICY FIRST (overrides router-level)
|
||||
|
|
|
|||
|
|
@ -93,4 +93,5 @@ async def async_raise_no_deployment_exception(
|
|||
cooldown_time=_cooldown_time,
|
||||
enable_pre_call_checks=litellm_router_instance.enable_pre_call_checks,
|
||||
cooldown_list=cooldown_list_ids,
|
||||
model_ids=model_ids,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -99,23 +99,13 @@ def setup(
|
|||
|
||||
def check_limits(kwargs: Mapping[str, object]) -> None:
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit
|
||||
|
||||
current_cost: Final = litellm._current_cost # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor
|
||||
if litellm.max_budget and current_cost > litellm.max_budget:
|
||||
raise litellm.BudgetExceededError(current_cost=current_cost, max_budget=litellm.max_budget)
|
||||
metadata: Final = kwargs.get("metadata")
|
||||
if isinstance(metadata, Mapping):
|
||||
typed_metadata: Final = cast( # cast-ok: runtime Mapping check establishes read-only metadata
|
||||
Mapping[str, object], metadata
|
||||
)
|
||||
previous: Final = typed_metadata.get("previous_models")
|
||||
if (
|
||||
isinstance(previous, list)
|
||||
and litellm.num_retries_per_request is not None
|
||||
and len(cast(list[object], previous)) # cast-ok: runtime list check establishes the retry history
|
||||
>= litellm.num_retries_per_request
|
||||
):
|
||||
raise RuntimeError("Max retries per request hit!")
|
||||
if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request):
|
||||
raise RuntimeError("Max retries per request hit!")
|
||||
|
||||
|
||||
def finalize(
|
||||
|
|
|
|||
|
|
@ -75,3 +75,9 @@ class HiddenParams(OpenAIObject):
|
|||
data: Final = super().model_dump(**kwargs)
|
||||
data["_response_ms"] = self._response_ms
|
||||
return data
|
||||
|
||||
|
||||
class CachedTokensDetails(BaseModel):
|
||||
text_tokens: int | None = None
|
||||
audio_tokens: int | None = None
|
||||
image_tokens: int | None = None
|
||||
|
|
|
|||
|
|
@ -91,6 +91,8 @@ from litellm.types.responses.main import (
|
|||
OutputImageGenerationCall,
|
||||
)
|
||||
|
||||
from .base import CachedTokensDetails
|
||||
|
||||
FileContent = IO[bytes] | bytes | PathLike
|
||||
|
||||
FileTypes = (
|
||||
|
|
@ -1288,6 +1290,7 @@ class OutputTokensDetails(BaseLiteLLMOpenAIResponseObject):
|
|||
class InputTokensDetails(BaseLiteLLMOpenAIResponseObject):
|
||||
audio_tokens: int | None = None
|
||||
cached_tokens: int = 0
|
||||
cached_tokens_details: CachedTokensDetails | None = None
|
||||
text_tokens: int | None = None
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
|
@ -2254,10 +2257,17 @@ class OpenAIRealtimeInputAudioTranscriptionCompleted(TypedDict):
|
|||
usage: NotRequired[ReadOnly[Mapping[str, object]]]
|
||||
|
||||
|
||||
class OpenAIRealtimeCachedTokensDetails(TypedDict, total=False):
|
||||
text_tokens: ReadOnly[int]
|
||||
audio_tokens: ReadOnly[int]
|
||||
image_tokens: ReadOnly[int]
|
||||
|
||||
|
||||
class OpenAIRealtimeUsageTokenDetails(TypedDict):
|
||||
audio_tokens: ReadOnly[int]
|
||||
text_tokens: ReadOnly[int]
|
||||
cached_tokens: NotRequired[ReadOnly[int]]
|
||||
cached_tokens_details: NotRequired[ReadOnly[OpenAIRealtimeCachedTokensDetails]]
|
||||
|
||||
|
||||
class OpenAIRealtimeResponseUsage(TypedDict):
|
||||
|
|
|
|||
8
litellm/types/proxy/auth/auth_checks.py
Normal file
8
litellm/types/proxy/auth/auth_checks.py
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
"""Failure values raised by `litellm/proxy/auth/auth_checks.py`. Kept free of `litellm` imports so any proxy module can import them without joining the `litellm.proxy` import cycle."""
|
||||
|
||||
|
||||
class UserNotFoundError(ValueError):
|
||||
"""The user row is provably absent, as opposed to merely unreadable, so a caller that reads a missing row as no user-level limits can key on it without also swallowing a database that would not answer."""
|
||||
|
||||
def __init__(self, user_id: str) -> None:
|
||||
super().__init__(f"User doesn't exist in db. 'user_id'={user_id}. Create user via `/user/new` call.")
|
||||
|
|
@ -645,6 +645,7 @@ class RouterErrors(enum.Enum):
|
|||
|
||||
user_defined_ratelimit_error = "Deployment over user-defined ratelimit."
|
||||
no_deployments_available = "No deployments available for selected model"
|
||||
all_deployments_in_cooldown = "All deployments for selected model are in cooldown"
|
||||
no_deployments_with_tag_routing = "Not allowed to access model due to tags configuration"
|
||||
no_deployments_with_provider_budget_routing = "No deployments available - crossed budget"
|
||||
no_healthy_deployments = "There are no healthy deployments for this model"
|
||||
|
|
@ -868,6 +869,11 @@ class RouterRateLimitErrorBasic(ValueError):
|
|||
super().__init__(_message)
|
||||
|
||||
|
||||
class RouterErrorTypes(str, enum.Enum):
|
||||
rate_limit_error = "rate_limit_error"
|
||||
all_deployments_in_cooldown = "all_deployments_in_cooldown"
|
||||
|
||||
|
||||
class RouterRateLimitError(ValueError):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -875,12 +881,25 @@ class RouterRateLimitError(ValueError):
|
|||
cooldown_time: float,
|
||||
enable_pre_call_checks: bool,
|
||||
cooldown_list: list,
|
||||
model_ids: Sequence[str] = (),
|
||||
) -> None:
|
||||
self.model = model
|
||||
self.cooldown_time = cooldown_time
|
||||
self.enable_pre_call_checks = enable_pre_call_checks
|
||||
self.cooldown_list = cooldown_list
|
||||
_message = f"{RouterErrors.no_deployments_available.value}, Try again in {cooldown_time} seconds. Passed model={model}. pre-call-checks={enable_pre_call_checks}, cooldown_list={cooldown_list}"
|
||||
self.all_deployments_in_cooldown = bool(model_ids) and frozenset(model_ids) <= frozenset(cooldown_list)
|
||||
self.type = (
|
||||
RouterErrorTypes.all_deployments_in_cooldown.value
|
||||
if self.all_deployments_in_cooldown
|
||||
else RouterErrorTypes.rate_limit_error.value
|
||||
)
|
||||
_reason: Final = (
|
||||
f" {RouterErrors.all_deployments_in_cooldown.value}." if self.all_deployments_in_cooldown else ""
|
||||
)
|
||||
_message: Final = (
|
||||
f"{RouterErrors.no_deployments_available.value}, Try again in {cooldown_time} seconds.{_reason} "
|
||||
f"Passed model={model}. pre-call-checks={enable_pre_call_checks}, cooldown_list={cooldown_list}"
|
||||
)
|
||||
super().__init__(_message)
|
||||
|
||||
|
||||
|
|
@ -889,6 +908,14 @@ class RouterModelGroupAliasItem(TypedDict):
|
|||
hidden: bool # if 'True', don't return on `.get_model_list`
|
||||
|
||||
|
||||
class RetryAttemptRecord(TypedDict):
|
||||
model_group: ReadOnly[str | None]
|
||||
deployment_id: ReadOnly[str | None]
|
||||
exception_type: ReadOnly[str]
|
||||
exception_string: ReadOnly[str]
|
||||
attempted_retries: ReadOnly[int | None]
|
||||
|
||||
|
||||
VALID_LITELLM_ENVIRONMENTS = [
|
||||
"development",
|
||||
"staging",
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ from litellm._logging import verbose_logger
|
|||
from litellm._uuid import uuid
|
||||
from litellm.types.llms.base import (
|
||||
BaseLiteLLMOpenAIResponseObject,
|
||||
CachedTokensDetails,
|
||||
LiteLLMPydanticObjectBase,
|
||||
)
|
||||
from litellm.types.mcp import MCPServerCostInfo
|
||||
|
|
@ -252,6 +253,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
cache_creation_input_token_cost_priority: float | None # OpenAI priority service tier pricing
|
||||
cache_creation_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing
|
||||
cache_read_input_token_cost: float | None
|
||||
cache_read_input_audio_token_cost: ReadOnly[float | None]
|
||||
cache_read_input_token_cost_flex: float | None # OpenAI flex service tier pricing
|
||||
cache_read_input_token_cost_priority: float | None # OpenAI priority service tier pricing
|
||||
cache_read_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing
|
||||
|
|
@ -1710,6 +1712,9 @@ class PromptTokensDetailsWrapper(
|
|||
cache_creation_token_details: CacheCreationTokenDetails | None = None
|
||||
"""Details of cache creation tokens sent to the model. Used for tracking 5m/1h cache creation tokens for Anthropic prompt caching."""
|
||||
|
||||
cached_tokens_details: CachedTokensDetails | None = None
|
||||
"""Details of cached (cache-hit) tokens sent to the model. OpenAI realtime naming; carries the per-modality cache-read split."""
|
||||
|
||||
def __setattr__(self, name: str, value: object) -> None:
|
||||
super().__setattr__(name, value)
|
||||
if name == "cache_write_tokens":
|
||||
|
|
@ -1756,6 +1761,8 @@ class PromptTokensDetailsWrapper(
|
|||
del self.cache_creation_tokens
|
||||
if self.cache_creation_token_details is None:
|
||||
del self.cache_creation_token_details
|
||||
if self.cached_tokens_details is None:
|
||||
del self.cached_tokens_details
|
||||
|
||||
|
||||
class ServerToolUse(BaseModel):
|
||||
|
|
|
|||
|
|
@ -84,6 +84,7 @@ from litellm.constants import (
|
|||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
from litellm.litellm_core_utils.fallback_generalizations import (
|
||||
match_capability_generalizations,
|
||||
match_fill_missing_generalizations,
|
||||
)
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import redact_credentials_in_payload
|
||||
|
||||
|
|
@ -254,6 +255,7 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
_CALL_TYPE_ENUM_MAP: Final[dict] = {ct.value: ct for ct in CallTypes}
|
||||
_BACKFILL_MODES: Final = frozenset({"chat", "responses"})
|
||||
|
||||
# +-----------------------------------------------+
|
||||
# | |
|
||||
|
|
@ -1260,15 +1262,6 @@ async def _client_async_logging_helper(
|
|||
async_coroutine=logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
|
||||
)
|
||||
|
||||
################################################
|
||||
# Sync Logging Worker
|
||||
################################################
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]:
|
||||
"""
|
||||
|
|
@ -1500,6 +1493,8 @@ def post_call_processing(
|
|||
|
||||
|
||||
def client(original_function):
|
||||
from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit
|
||||
|
||||
Rules: Final = litellm_utils.Rules
|
||||
rules_obj: Final = Rules()
|
||||
|
||||
|
|
@ -1510,12 +1505,8 @@ def client(original_function):
|
|||
call_type = original_function.__name__
|
||||
if _is_async_request(kwargs):
|
||||
# [OPTIONAL] CHECK MAX RETRIES / REQUEST
|
||||
if litellm.num_retries_per_request is not None:
|
||||
# check if previous_models passed in as ['litellm_params']['metadata]['previous_models']
|
||||
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
|
||||
if previous_models is not None:
|
||||
if litellm.num_retries_per_request <= len(previous_models):
|
||||
raise Exception("Max retries per request hit!")
|
||||
if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request):
|
||||
raise Exception("Max retries per request hit!")
|
||||
|
||||
# MODEL CALL
|
||||
result = original_function(*args, **kwargs)
|
||||
|
|
@ -1574,12 +1565,8 @@ def client(original_function):
|
|||
)
|
||||
|
||||
# [OPTIONAL] CHECK MAX RETRIES / REQUEST
|
||||
if litellm.num_retries_per_request is not None:
|
||||
# check if previous_models passed in as ['litellm_params']['metadata]['previous_models']
|
||||
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
|
||||
if previous_models is not None:
|
||||
if litellm.num_retries_per_request <= len(previous_models):
|
||||
raise Exception("Max retries per request hit!")
|
||||
if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request):
|
||||
raise Exception("Max retries per request hit!")
|
||||
|
||||
# [OPTIONAL] CHECK CACHE
|
||||
print_verbose(
|
||||
|
|
@ -5819,6 +5806,14 @@ def _get_model_info_helper(
|
|||
):
|
||||
_model_info = None
|
||||
|
||||
if _model_info is not None and key is not None and _model_info.get("mode", "chat") in _BACKFILL_MODES:
|
||||
fill_missing: Final = match_fill_missing_generalizations(key, _model_info.get("litellm_provider", ""))
|
||||
if fill_missing is not None:
|
||||
_model_info = {
|
||||
**{k: v for k, v in fill_missing.items() if k not in _model_info},
|
||||
**_model_info,
|
||||
}
|
||||
|
||||
if _model_info is None:
|
||||
generalization: Final = _get_model_info_from_generalization(
|
||||
model=model,
|
||||
|
|
@ -5882,6 +5877,7 @@ def _get_model_info_helper(
|
|||
"cache_creation_input_token_cost_ultrafast", None
|
||||
),
|
||||
cache_read_input_token_cost=_model_info.get("cache_read_input_token_cost", None),
|
||||
cache_read_input_audio_token_cost=_model_info.get("cache_read_input_audio_token_cost", None),
|
||||
prompt_cache_min_tokens=_model_info.get("prompt_cache_min_tokens", None),
|
||||
cache_read_input_token_cost_above_200k_tokens=_model_info.get(
|
||||
"cache_read_input_token_cost_above_200k_tokens", None
|
||||
|
|
|
|||
|
|
@ -5513,7 +5513,8 @@
|
|||
},
|
||||
"azure/gpt-realtime-2025-08-28": {
|
||||
"cache_creation_input_audio_token_cost": 4e-06,
|
||||
"cache_read_input_token_cost": 4e-06,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-03-02",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
|
|
@ -5546,7 +5547,8 @@
|
|||
},
|
||||
"azure/gpt-realtime-1.5-2026-02-23": {
|
||||
"cache_creation_input_audio_token_cost": 4e-06,
|
||||
"cache_read_input_token_cost": 4e-06,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-08-24",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
|
|
@ -5683,6 +5685,7 @@
|
|||
},
|
||||
"azure/gpt-realtime-mini": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-07,
|
||||
|
|
@ -5715,6 +5718,7 @@
|
|||
},
|
||||
"azure/gpt-realtime-mini-2025-10-06": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-07,
|
||||
|
|
@ -7409,6 +7413,80 @@
|
|||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/gpt-chat-latest": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"deprecation_date": "2026-12-02",
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"medium"
|
||||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": 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
|
||||
},
|
||||
"azure/chat-latest": {
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"deprecation_date": "2026-12-02",
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"medium"
|
||||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": 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
|
||||
},
|
||||
"azure/us/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
@ -7675,6 +7753,43 @@
|
|||
"supports_web_search": true,
|
||||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"azure/us/gpt-chat-latest": {
|
||||
"cache_read_input_token_cost": 5.5e-07,
|
||||
"deprecation_date": "2026-12-02",
|
||||
"input_cost_per_token": 5.5e-06,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 272000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.3e-05,
|
||||
"reasoning_effort_levels": [
|
||||
"medium"
|
||||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/",
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": 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
|
||||
},
|
||||
"azure/eu/gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.875e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.375e-05,
|
||||
|
|
@ -23053,34 +23168,6 @@
|
|||
"output_cost_per_token": 0.0,
|
||||
"source": "https://fireworks.ai/pricing"
|
||||
},
|
||||
"friendliai/meta-llama-3.1-70b-instruct": {
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "friendliai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6e-07,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"friendliai/meta-llama-3.1-8b-instruct": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "friendliai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1e-07,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"friendliai/zai-org/GLM-5.3-Flash": {
|
||||
"litellm_provider": "friendliai",
|
||||
"max_input_tokens": 1048576,
|
||||
|
|
@ -32678,6 +32765,7 @@
|
|||
},
|
||||
"gpt-realtime": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-01-20",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
|
|
@ -32711,6 +32799,7 @@
|
|||
},
|
||||
"gpt-realtime-1.5": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
"input_cost_per_image_token": 5e-06,
|
||||
|
|
@ -32847,6 +32936,7 @@
|
|||
"gpt-realtime-mini": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"deprecation_date": "2027-01-20",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_token": 6e-07,
|
||||
|
|
@ -32878,6 +32968,7 @@
|
|||
},
|
||||
"gpt-realtime-2025-08-28": {
|
||||
"cache_creation_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_audio_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"deprecation_date": "2027-01-20",
|
||||
"input_cost_per_audio_token": 3.2e-05,
|
||||
|
|
@ -57625,8 +57716,9 @@
|
|||
},
|
||||
{
|
||||
"name": "claude-adaptive-thinking",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"description": "Claude at version 4.6 or higher, in any id shape that contains claude-<family>-: minors 4.6 through 4.99, any later major-minor, and bare 5+ majors so a new family shaped like claude-fable-5 matches. Requiring the claude- prefix keeps non-Claude names such as team-sonnet-5-1 out. The minor is capped at two digits so an 8-digit date suffix such as claude-opus-4-20250514 is never read as a >= 4.6 minor. Turns on adaptive thinking for new versions and new families with no code change.",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|[5-9](?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"fill_missing_for_providers": ["anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai-anthropic_models"],
|
||||
"description": "Claude at version 4.6 or higher, in any id shape that contains claude-<family>-: minors 4.6 through 4.99, any later major-minor, and bare 5+ majors so a new family shaped like claude-fable-5 matches. Requiring the claude- prefix keeps non-Claude names such as team-sonnet-5-1 out. The minor is capped at two digits so an 8-digit date suffix such as claude-opus-4-20250514 is never read as a >= 4.6 minor. Two-digit majors are deliberately not matched so ids like claude-opus-41 (4.1) are not read as major 41. Turns on adaptive thinking for new versions and new families with no code change.",
|
||||
"model_info": {
|
||||
"supports_adaptive_thinking": true
|
||||
}
|
||||
|
|
@ -57634,6 +57726,7 @@
|
|||
{
|
||||
"name": "claude-legacy-thinking",
|
||||
"pattern": "claude-[a-z]+-4[-._]6(?!\\d)",
|
||||
"fill_missing_for_providers": ["anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai-anthropic_models"],
|
||||
"description": "Claude at version 4.6 exactly, in any id shape that contains claude-<family>-4-6 (dotted and underscored minors included, dated releases such as claude-sonnet-4-6-20260219 too). The 4.6 family is adaptive-thinking yet still accepts legacy thinking.type=enabled with budget_tokens, so the caller's hard budget cap is forwarded verbatim instead of being rewritten to an uncapped output_config.effort. The lookahead keeps two-digit minors such as 4-60 from matching. 4.7+ and 5+ majors reject the legacy shape and stay on the adaptive translation.",
|
||||
"model_info": {
|
||||
"supports_legacy_thinking": true
|
||||
|
|
@ -57649,8 +57742,9 @@
|
|||
},
|
||||
{
|
||||
"name": "claude-mid-conversation-system",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"description": "Claude at version 4.8 or higher, in any id shape that contains claude-<family>-: minors 4.8 through 4.99, any later major-minor, and bare 5+ majors so a new family like claude-fable-5 matches. Anthropic introduced mid-conversation system messages with Opus 4.8 and every newer Claude keeps them; 4.7 and below reject the system role inside messages.",
|
||||
"pattern": "claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|[5-9](?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
|
||||
"fill_missing_for_providers": ["anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai-anthropic_models"],
|
||||
"description": "Claude at version 4.8 or higher, in any id shape that contains claude-<family>-: minors 4.8 through 4.99, any later major-minor, and bare 5+ majors so a new family like claude-fable-5 matches. Two-digit majors are deliberately not matched so ids like claude-opus-41 (4.1) are not read as major 41. Anthropic introduced mid-conversation system messages with Opus 4.8 and every newer Claude keeps them; 4.7 and below reject the system role inside messages.",
|
||||
"model_info": {
|
||||
"supports_mid_conversation_system": true
|
||||
}
|
||||
|
|
@ -57666,6 +57760,7 @@
|
|||
{
|
||||
"name": "openai-reasoning-family-baseline",
|
||||
"pattern": "^(?!.*search-api)(?:[a-z0-9_.-]+/)*(?:ft:)?(?:o[1-9]\\d*(?![a-z0-9])|gpt-[5-9](?:\\.\\d+)?(?![0-9.])|(?:gpt-\\d+(?:\\.\\d+)?(?:-[a-z0-9]+)*-)?(?:codex|deep-research|chat-latest)(?![a-z0-9]))",
|
||||
"fill_missing_for_providers": ["azure", "azure_ai", "openai"],
|
||||
"description": "OpenAI reasoning families by id shape, under any provider namespace and with an optional ft: prefix: the o-series (o1, o3-pro, o4-mini), gpt-5 through gpt-9 majors including dotted minors and suffixed variants (gpt-5.5-cyber, gpt-6-astra), and the codex, deep-research and chat-latest lines when standalone or on a gpt base. gpt-5-search-api is excluded because it is a search-only surface. Every model here is a reasoning model, and the Responses API drops the caller's reasoning param for any mapped OpenAI model whose info lacks supports_reasoning, so an id the registry has not named yet keeps its reasoning settings instead of silently losing them. Rules lose to exact entries. Carries no mode and no pricing, so cost stays on the standard unpriced behavior.",
|
||||
"model_info": {
|
||||
"supports_reasoning": true
|
||||
|
|
|
|||
|
|
@ -5,5 +5,5 @@ reason = "diskcache has no fixed release published; remove this entry once one e
|
|||
|
||||
[[IgnoredVulns]]
|
||||
id = "GHSA-h7x2-h6g9-p789"
|
||||
ignoreUntil = 2026-09-14
|
||||
reason = "mlflow has no fixed release published; remove this entry once one exists"
|
||||
ignoreUntil = 2026-10-14
|
||||
reason = "mlflow has no fixed release published (3.16.0, 2026-09-04, and master still store gateway secret api_base unvalidated); remove this entry once one exists"
|
||||
|
|
|
|||
|
|
@ -22,6 +22,11 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|||
pytestmark = pytest.mark.asyncio(loop_scope="session")
|
||||
|
||||
|
||||
def _frozen_cache() -> UserApiKeyCache:
|
||||
"""The org entries carry a 5s TTL; a frozen clock keeps a slow first call from expiring them mid-test."""
|
||||
return UserApiKeyCache(in_memory_cache=InMemoryCache(clock=lambda: 1_000_000.0), redis_cache=None)
|
||||
|
||||
|
||||
def _dead_db() -> MagicMock:
|
||||
prisma = MagicMock(name="prisma_client")
|
||||
prisma.db.query_first = AsyncMock(return_value=None)
|
||||
|
|
@ -58,9 +63,10 @@ async def test_join_binds_the_membership_to_the_requested_team(prisma):
|
|||
data={"user_id": user_id, "team_id": team_b, "litellm_budget_table": {"connect": {"budget_id": f"b-{run}"}}}
|
||||
)
|
||||
|
||||
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(), redis_cache=None)
|
||||
cache = _frozen_cache()
|
||||
refs = AuthObjectRefs(user_id=user_id, team_id=team_a, membership_user_id=user_id, organization_id=org_id)
|
||||
await prefetch_auth_objects(refs=refs, user_api_key_cache=cache, prisma_client=prisma)
|
||||
assert cache.in_memory_cache.get_cache(f"org_id:{org_id}") is not None
|
||||
|
||||
dead_db = _dead_db()
|
||||
membership = await get_team_membership(
|
||||
|
|
@ -100,7 +106,7 @@ async def test_join_reads_team_model_aliases_from_the_mapped_column(prisma):
|
|||
where={"team_id": team_id}, include={"litellm_model_table": True}
|
||||
)
|
||||
|
||||
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(), redis_cache=None)
|
||||
cache = _frozen_cache()
|
||||
refs = AuthObjectRefs(user_id=None, team_id=team_id, membership_user_id=None, organization_id=None)
|
||||
await prefetch_auth_objects(refs=refs, user_api_key_cache=cache, prisma_client=prisma)
|
||||
|
||||
|
|
@ -144,7 +150,7 @@ async def test_join_reads_null_nested_lists_the_way_prisma_does(prisma):
|
|||
where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, include={"litellm_budget_table": True}
|
||||
)
|
||||
|
||||
cache = UserApiKeyCache(in_memory_cache=InMemoryCache(), redis_cache=None)
|
||||
cache = _frozen_cache()
|
||||
refs = AuthObjectRefs(user_id=user_id, team_id=team_id, membership_user_id=user_id, organization_id=None)
|
||||
await prefetch_auth_objects(refs=refs, user_api_key_cache=cache, prisma_client=prisma)
|
||||
|
||||
|
|
|
|||
|
|
@ -2139,7 +2139,7 @@ async def test_model_info_alias_without_prisma(hidden):
|
|||
user_api_key_dict=UserAPIKeyAuth(models=[]),
|
||||
)
|
||||
|
||||
models = resp["data"]
|
||||
models = json.loads(resp.body)["data"]
|
||||
|
||||
alias_found = any(
|
||||
m["model_name"] == model_alias
|
||||
|
|
@ -2203,7 +2203,7 @@ async def test_proxy_model_group_alias_checks(prisma_client, hidden): # noqa: F
|
|||
resp = await model_info_v1(
|
||||
user_api_key_dict=UserAPIKeyAuth(models=[]),
|
||||
)
|
||||
models = resp["data"]
|
||||
models = json.loads(resp.body)["data"]
|
||||
is_model_alias_in_list = False
|
||||
for item in models:
|
||||
if model_alias == item["model_name"]:
|
||||
|
|
@ -2280,7 +2280,7 @@ async def test_proxy_model_group_info_rerank(prisma_client): # noqa: F811 # py
|
|||
resp = await model_info_v1(
|
||||
user_api_key_dict=UserAPIKeyAuth(models=[]),
|
||||
)
|
||||
models = resp["data"]
|
||||
models = json.loads(resp.body)["data"]
|
||||
assert models[0]["model_info"]["mode"] == "rerank"
|
||||
resp = await model_group_info(
|
||||
user_api_key_dict=UserAPIKeyAuth(models=[]),
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import json
|
||||
import os
|
||||
import traceback
|
||||
from dotenv import load_dotenv
|
||||
|
|
@ -628,17 +629,29 @@ def test_deployment_callback_respects_cooldown_time(model_list):
|
|||
assert mock_set.call_args.kwargs["time_to_cooldown"] == 0
|
||||
|
||||
|
||||
def test_log_retry(model_list):
|
||||
"""Test if the '_log_retry' function is working correctly"""
|
||||
import time
|
||||
|
||||
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
|
||||
def test_log_retry(model_list, metadata_key):
|
||||
"""log_retry appends one flat record per failed attempt and copies neither the request kwargs nor
|
||||
the request metadata into it"""
|
||||
router = Router(model_list=model_list)
|
||||
new_kwargs = router.log_retry(
|
||||
kwargs={"metadata": {}},
|
||||
e=Exception(),
|
||||
kwargs={
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": "sk-must-not-be-recorded",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
metadata_key: {"model_info": {"id": "deployment-1"}, "attempted_retries": 2, "user_api_key": "sk-proxy"},
|
||||
},
|
||||
e=litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo"),
|
||||
)
|
||||
assert "metadata" in new_kwargs
|
||||
assert "previous_models" in new_kwargs["metadata"]
|
||||
assert json.loads(json.dumps(new_kwargs[metadata_key]["previous_models"])) == [
|
||||
{
|
||||
"model_group": "gpt-3.5-turbo",
|
||||
"deployment_id": "deployment-1",
|
||||
"exception_type": "RateLimitError",
|
||||
"exception_string": "litellm.RateLimitError: slow down",
|
||||
"attempted_retries": 2,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_update_usage(model_list):
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ tests/rust-python-harness/
|
|||
- Examples: `run e2e_parity --surface sdk --function ocr`, `run unit_tests_parity --function ocr --pytest-arg=-x`, or `run all --function ocr`
|
||||
- `cli/catalog.py` discovers strategies, validates their Python definitions, and orders them; `cli/__init__.py` builds the Click command tree; `cli/commands.py` runs selected cases
|
||||
- `e2e_parity/` compares SDK objects, exceptions, callbacks, and streams, or gateway HTTP responses
|
||||
- `trace_parity/` compares mapped operations, call counts, and required execution ordering; before running it rebuilds the native bridge with the `trace-parity` feature whenever `litellm-rust` sources are newer than the installed extension (`shared/native_build.py`)
|
||||
- `trace_parity/` prints every collected Python call under `litellm/` and every Rust span without comparing them; mappings only filter the separate unit-test mapping strategy. Before running it rebuilds the native bridge with the `trace-parity` feature whenever `litellm-rust` sources are newer than the installed extension (`shared/native_build.py`)
|
||||
- E2E and trace strategies load their registered module cases and run surface-specific execution from their folders
|
||||
- `unit_tests_mapping/contracts.py` owns typed harness-side mapping contracts, per-function contracts live below `cases/`, and `mappings.py` exports the registry; live test discovery derives unmapped Python and Rust-only tests without an exhaustive manifest
|
||||
- `unit_tests_mapping/runner.py` validates confirmed mappings against the live Python and Rust inventories and attaches the derived status report
|
||||
|
|
|
|||
|
|
@ -58,16 +58,29 @@ def _strategy_command(strategy: Strategy) -> click.Command:
|
|||
help=runner_argument.help,
|
||||
)
|
||||
)
|
||||
for runner_option in strategy.definition.runner_options:
|
||||
name: Final = runner_option.option.removeprefix("--").replace("-", "_")
|
||||
params.append(
|
||||
click.Option(
|
||||
(runner_option.option, name),
|
||||
type=click.Choice(runner_option.choices),
|
||||
help=runner_option.help,
|
||||
)
|
||||
)
|
||||
|
||||
def run_strategy(
|
||||
sdk_functions: tuple[str, ...],
|
||||
surface: str | None = None,
|
||||
runner_args: tuple[str, ...] = (),
|
||||
**runner_options: str | None,
|
||||
) -> int:
|
||||
selected_functions: Final = cast(frozenset[SdkFunction], frozenset(sdk_functions))
|
||||
selected_surface: Final = cast(Surface | None, surface)
|
||||
cases: Final = select_cases((strategy,), selected_functions, selected_surface)
|
||||
return run_command((strategy,), cases, runner_args)
|
||||
option_args: Final = tuple(
|
||||
f"--{name.replace('_', '-')}={value}" for name, value in runner_options.items() if value is not None
|
||||
)
|
||||
return run_command((strategy,), cases, (*runner_args, *option_args))
|
||||
|
||||
return click.Command(
|
||||
strategy.id,
|
||||
|
|
|
|||
|
|
@ -21,9 +21,7 @@ def _load_strategy_module(name: str, folder: Path, prefix: str | None) -> Module
|
|||
if prefix is not None:
|
||||
return importlib.import_module(f"{prefix}.{name}")
|
||||
module_name: Final = _synthetic_module_name(folder)
|
||||
spec: Final = importlib.util.spec_from_file_location(
|
||||
module_name, folder / "__init__.py"
|
||||
)
|
||||
spec: Final = importlib.util.spec_from_file_location(module_name, folder / "__init__.py")
|
||||
if spec is None or spec.loader is None:
|
||||
raise ValueError(f"{folder}: cannot load strategy package")
|
||||
module: Final = importlib.util.module_from_spec(spec)
|
||||
|
|
@ -59,9 +57,7 @@ def _load_strategy(name: str, folder: Path, prefix: str | None) -> Strategy:
|
|||
if duplicates:
|
||||
raise ValueError(f"{folder}: duplicate strategy cases: {duplicates}")
|
||||
expected: Final = frozenset(
|
||||
(surface, function)
|
||||
for surface in (definition.surfaces or (None,))
|
||||
for function in SDK_FUNCTIONS
|
||||
(surface, function) for surface in (definition.surfaces or (None,)) for function in SDK_FUNCTIONS
|
||||
)
|
||||
actual: Final = frozenset(keys)
|
||||
if actual != expected:
|
||||
|
|
@ -73,8 +69,7 @@ def _load_strategy(name: str, folder: Path, prefix: str | None) -> Strategy:
|
|||
incompatible: Final = tuple(
|
||||
(case.surface, case.sdk_function)
|
||||
for case in definition.cases
|
||||
if case.spec.disposition is CaseDisposition.RUNNABLE
|
||||
and not isinstance(case.spec, definition.runnable_spec)
|
||||
if case.spec.disposition is CaseDisposition.RUNNABLE and not isinstance(case.spec, definition.runnable_spec)
|
||||
)
|
||||
if incompatible:
|
||||
raise ValueError(f"{folder}: runnable cases do not match {definition.runnable_spec.__name__}: {incompatible}")
|
||||
|
|
@ -102,14 +97,10 @@ def _load_strategy(name: str, folder: Path, prefix: str | None) -> Strategy:
|
|||
def load_catalog(root: Path | None = None) -> tuple[Strategy, ...]:
|
||||
resolved: Final = STRATEGIES_ROOT if root is None else root
|
||||
prefix: Final = _STRATEGIES_PACKAGE.__name__ if resolved == STRATEGIES_ROOT else None
|
||||
folders: Final = tuple(
|
||||
info.name for info in pkgutil.iter_modules([str(resolved)]) if info.ispkg
|
||||
)
|
||||
folders: Final = tuple(info.name for info in pkgutil.iter_modules([str(resolved)]) if info.ispkg)
|
||||
if not folders:
|
||||
raise ValueError(f"No strategy packages found below {resolved}")
|
||||
strategies: Final = tuple(
|
||||
_load_strategy(name, resolved / name, prefix) for name in sorted(folders)
|
||||
)
|
||||
strategies: Final = tuple(_load_strategy(name, resolved / name, prefix) for name in sorted(folders))
|
||||
ids: Final = [strategy.id for strategy in strategies]
|
||||
if len(set(ids)) != len(ids):
|
||||
raise ValueError(f"Duplicate strategy id in {resolved}")
|
||||
|
|
|
|||
|
|
@ -21,8 +21,7 @@ def select_cases(
|
|||
case
|
||||
for strategy in strategies
|
||||
for case in strategy.cases
|
||||
if (not sdk_functions or case.sdk_function in sdk_functions)
|
||||
and (surface is None or case.surface == surface)
|
||||
if (not sdk_functions or case.sdk_function in sdk_functions) and (surface is None or case.surface == surface)
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -32,8 +31,7 @@ def run_command(
|
|||
runner_args: Sequence[str] = (),
|
||||
) -> int:
|
||||
grouped: Final = {
|
||||
strategy.id: tuple(case for case in cases if case.strategy_id == strategy.id)
|
||||
for strategy in strategies
|
||||
strategy.id: tuple(case for case in cases if case.strategy_id == strategy.id) for strategy in strategies
|
||||
}
|
||||
visible: Final = tuple(strategy for strategy in strategies if grouped[strategy.id])
|
||||
runners: Final = tuple(replace(strategy, cases=grouped[strategy.id]) for strategy in visible)
|
||||
|
|
|
|||
|
|
@ -242,7 +242,7 @@ def _assert_unavailable_cell(strategy: Strategy, case: HarnessCase, section_titl
|
|||
def test_every_unavailable_case_finishes_and_explains_itself() -> None:
|
||||
section_titles: Final = {
|
||||
"e2e_parity": "End-to-end parity outcomes",
|
||||
"trace_parity": "trace comparisons",
|
||||
"trace_parity": "traces",
|
||||
"unit_tests_mapping": "Python/Rust unit-test mappings",
|
||||
"unit_tests_parity": "Python backend parity outcomes",
|
||||
"unit_tests_rust": "Native Rust unit-test outcomes",
|
||||
|
|
@ -359,6 +359,25 @@ def test_strategy_command_forwards_repeated_filters_and_runner_arguments(
|
|||
]
|
||||
|
||||
|
||||
def test_trace_command_forwards_engine_and_scenario(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
cli: Final = importlib.import_module("tests.rust-python-harness.cli")
|
||||
captured: list[tuple[str, ...]] = []
|
||||
|
||||
def capture_run(
|
||||
strategies: Sequence[Strategy],
|
||||
cases: Sequence[HarnessCase],
|
||||
runner_args: Sequence[str] = (),
|
||||
) -> int:
|
||||
del strategies, cases
|
||||
captured.append(tuple(runner_args))
|
||||
return 0
|
||||
|
||||
monkeypatch.setattr(cli, "run_command", capture_run)
|
||||
|
||||
assert main(["run", "trace_parity", "--scenario", "async-mistral", "--engine", "python"]) == 0
|
||||
assert captured == [("async-mistral", "--engine=python")]
|
||||
|
||||
|
||||
def test_omitted_surface_selects_every_strategy_surface(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
cli: Final = importlib.import_module("tests.rust-python-harness.cli")
|
||||
selected: list[str] = []
|
||||
|
|
|
|||
|
|
@ -19,9 +19,7 @@ def subprocess_test_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
def cargo_project(tmp_path: Path) -> Callable[[str, str], Path]:
|
||||
def create(package: str, source: str) -> Path:
|
||||
manifest: Final = tmp_path / "Cargo.toml"
|
||||
manifest.write_text(
|
||||
f'[package]\nname = "{package}"\nversion = "0.1.0"\nedition = "2021"\n[workspace]\n'
|
||||
)
|
||||
manifest.write_text(f'[package]\nname = "{package}"\nversion = "0.1.0"\nedition = "2021"\n[workspace]\n')
|
||||
(tmp_path / "src").mkdir()
|
||||
(tmp_path / "src/lib.rs").write_text(source)
|
||||
return manifest
|
||||
|
|
|
|||
|
|
@ -16,6 +16,11 @@ _RUST_ROOT: Final = "litellm-rust"
|
|||
_LOCKFILE: Final = "Cargo.lock"
|
||||
_SOURCE_SUFFIXES: Final = frozenset({".rs", ".toml"})
|
||||
_FAILURE_OUTPUT_LINES: Final = 15
|
||||
_TRACE_CHECK: Final = (
|
||||
"from litellm.rust_bridge import get_native_bridge; "
|
||||
"bridge = get_native_bridge(); "
|
||||
"raise SystemExit(0 if bridge is not None and getattr(bridge, '_trace', None) is not None else 1)"
|
||||
)
|
||||
|
||||
|
||||
def needs_rebuild(native_mtime: float | None, newest_source_mtime: float | None) -> bool:
|
||||
|
|
@ -73,6 +78,17 @@ def _rebuild(repo_root: Path) -> tuple[bool, str]:
|
|||
return completed.returncode == 0, "\n".join(lines[-_FAILURE_OUTPUT_LINES:])
|
||||
|
||||
|
||||
def _installed_bridge_has_trace(repo_root: Path) -> bool:
|
||||
completed: Final = subprocess.run(
|
||||
(sys.executable, "-c", _TRACE_CHECK),
|
||||
cwd=repo_root,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
check=False,
|
||||
)
|
||||
return completed.returncode == 0
|
||||
|
||||
|
||||
def trace_bridge_error() -> str | None:
|
||||
bridge: Final = get_native_bridge()
|
||||
if bridge is None:
|
||||
|
|
@ -85,7 +101,10 @@ def trace_bridge_error() -> str | None:
|
|||
def ensure_trace_bridge(repo_root: Path) -> str | None:
|
||||
native_path: Final = _native_module_path()
|
||||
native_mtime: Final = native_path.stat().st_mtime if native_path is not None and native_path.exists() else None
|
||||
if needs_rebuild(native_mtime, _newest_source_mtime(repo_root)):
|
||||
rebuild_required: Final = needs_rebuild(
|
||||
native_mtime, _newest_source_mtime(repo_root)
|
||||
) or not _installed_bridge_has_trace(repo_root)
|
||||
if rebuild_required:
|
||||
print(f"Rebuilding native Rust bridge ({BRIDGE_FEATURE} feature)...", flush=True)
|
||||
succeeded: Final
|
||||
output: Final
|
||||
|
|
|
|||
|
|
@ -191,6 +191,7 @@ class _RecordingHandler(LocalHttpHandler):
|
|||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
|
||||
def _recording_provider(spec: UpstreamEndpoint) -> AbstractContextManager[_RecordingProvider]:
|
||||
return serve_in_thread(_RecordingProvider(spec))
|
||||
|
||||
|
|
|
|||
|
|
@ -15,6 +15,8 @@ from .cassette import deserialize_cassette, serialize_cassette
|
|||
from .recording import RecordedInteraction
|
||||
|
||||
FIXTURE_SCHEMA_VERSION: Final = 1
|
||||
|
||||
|
||||
class FixtureInput(Protocol):
|
||||
def canonical_input(self) -> dict[str, object]: ...
|
||||
|
||||
|
|
|
|||
|
|
@ -45,8 +45,8 @@ class _Upstream(LocalHttpServer):
|
|||
super().__init__(("127.0.0.1", 0), _UpstreamHandler)
|
||||
self.response_status: Final = status
|
||||
|
||||
class _UpstreamHandler(LocalHttpHandler):
|
||||
|
||||
class _UpstreamHandler(LocalHttpHandler):
|
||||
def do_POST(self) -> None:
|
||||
length: Final = int(self.headers.get("content-length") or "0")
|
||||
self.rfile.read(length)
|
||||
|
|
@ -59,6 +59,7 @@ class _UpstreamHandler(LocalHttpHandler):
|
|||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
|
||||
def _upstream(status: int = 200) -> AbstractContextManager[_Upstream]:
|
||||
return serve_in_thread(_Upstream(status))
|
||||
|
||||
|
|
|
|||
|
|
@ -238,6 +238,7 @@ class _ControlledUpstreamHandler(LocalHttpHandler):
|
|||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
|
||||
def _controlled_upstream(
|
||||
stream_chunks: tuple[bytes, ...] = _SSE_CHUNKS,
|
||||
) -> AbstractContextManager[_ControlledUpstream]:
|
||||
|
|
|
|||
|
|
@ -180,15 +180,11 @@ class HarnessRun:
|
|||
|
||||
@property
|
||||
def unique_checks(self) -> int:
|
||||
return len(
|
||||
{nodeid for result in self.results.values() for nodeid in result.collected}
|
||||
)
|
||||
return len({nodeid for result in self.results.values() for nodeid in result.collected})
|
||||
|
||||
@property
|
||||
def completed_checks(self) -> int:
|
||||
return len(
|
||||
{nodeid for result in self.results.values() for nodeid in result.completed}
|
||||
)
|
||||
return len({nodeid for result in self.results.values() for nodeid in result.completed})
|
||||
|
||||
@classmethod
|
||||
def from_cases(cls, cases: Iterable[HarnessCase]) -> HarnessRun:
|
||||
|
|
|
|||
|
|
@ -67,6 +67,13 @@ class RunnerArgumentDefinition:
|
|||
metavar: str = "ARG"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RunnerOptionDefinition:
|
||||
option: str
|
||||
help: str
|
||||
choices: tuple[str, ...]
|
||||
|
||||
|
||||
class StrategyRunner(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
|
|
@ -90,3 +97,4 @@ class StrategyDefinition:
|
|||
render: StrategyRenderer
|
||||
surfaces: tuple[Surface, ...] = ()
|
||||
runner_argument: RunnerArgumentDefinition | None = None
|
||||
runner_options: tuple[RunnerOptionDefinition, ...] = ()
|
||||
|
|
|
|||
|
|
@ -89,8 +89,8 @@ def test_ensure_trace_bridge_reports_failed_rebuild(tmp_path: Final, monkeypatch
|
|||
assert "boom" in message
|
||||
|
||||
|
||||
def test_ensure_trace_bridge_flags_missing_trace_feature_without_rebuild(
|
||||
tmp_path: Final, monkeypatch: pytest.MonkeyPatch
|
||||
def test_ensure_trace_bridge_rebuilds_when_trace_feature_is_missing(
|
||||
tmp_path: Final, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str]
|
||||
) -> None:
|
||||
native: Final = tmp_path / "_native.abi3.so"
|
||||
native.write_bytes(b"")
|
||||
|
|
@ -105,12 +105,17 @@ def test_ensure_trace_bridge_flags_missing_trace_feature_without_rebuild(
|
|||
state.rebuilt = True
|
||||
return True, ""
|
||||
|
||||
def fake_get_native_bridge() -> SimpleNamespace:
|
||||
assert state.rebuilt
|
||||
return SimpleNamespace(_trace=object())
|
||||
|
||||
monkeypatch.setattr(native_build, "_native_module_path", lambda: native)
|
||||
monkeypatch.setattr(native_build, "_rebuild", fake_rebuild)
|
||||
monkeypatch.setattr(native_build, "get_native_bridge", lambda: SimpleNamespace(_trace=None))
|
||||
monkeypatch.setattr(native_build, "_installed_bridge_has_trace", lambda repo_root: False)
|
||||
monkeypatch.setattr(native_build, "get_native_bridge", fake_get_native_bridge)
|
||||
|
||||
message: Final = native_build.ensure_trace_bridge(tmp_path)
|
||||
|
||||
assert message is not None
|
||||
assert "_trace" in message
|
||||
assert state.rebuilt is False
|
||||
assert message is None
|
||||
assert state.rebuilt is True
|
||||
assert "Rebuilding native Rust bridge" in capsys.readouterr().out
|
||||
|
|
|
|||
|
|
@ -2,7 +2,8 @@ from __future__ import annotations
|
|||
|
||||
import sys
|
||||
import threading
|
||||
from collections.abc import Generator, Iterator, Mapping
|
||||
import warnings
|
||||
from collections.abc import Callable, Generator, Iterator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
|
|
@ -32,6 +33,7 @@ class PythonProfiler:
|
|||
self._source_root: Final = str(source_root.resolve()) + "/"
|
||||
self._seen_frames: Final[set[FrameType]] = set()
|
||||
self._event_ids: Final[dict[FrameType, int]] = {}
|
||||
self._lock: Final = threading.Lock()
|
||||
self.events: Final[list[FunctionTraceEvent]] = []
|
||||
|
||||
def __call__(self, frame: FrameType, event: str, _arg: object) -> None:
|
||||
|
|
@ -40,14 +42,15 @@ class PythonProfiler:
|
|||
function_name: Final = self.function_name(frame)
|
||||
if function_name is None:
|
||||
return
|
||||
event_id: Final = len(self.events)
|
||||
parent_id: Final = next(
|
||||
(self._event_ids[ancestor] for ancestor in _frame_ancestors(frame) if ancestor in self._event_ids),
|
||||
None,
|
||||
)
|
||||
self._seen_frames.add(frame)
|
||||
self._event_ids[frame] = event_id
|
||||
self.events.append(FunctionTraceEvent(id=event_id, parent_id=parent_id, function=function_name))
|
||||
with self._lock:
|
||||
event_id: Final = len(self.events)
|
||||
parent_id: Final = next(
|
||||
(self._event_ids[ancestor] for ancestor in _frame_ancestors(frame) if ancestor in self._event_ids),
|
||||
None,
|
||||
)
|
||||
self._seen_frames.add(frame)
|
||||
self._event_ids[frame] = event_id
|
||||
self.events.append(FunctionTraceEvent(id=event_id, parent_id=parent_id, function=function_name))
|
||||
|
||||
def function_name(self, frame: FrameType) -> str | None:
|
||||
code: Final = frame.f_code
|
||||
|
|
@ -137,21 +140,51 @@ def _frame_ancestors(frame: FrameType) -> Generator[FrameType]:
|
|||
|
||||
|
||||
@contextmanager
|
||||
def profile_python(source_root: Path, *, threads: bool = False) -> Generator[PythonProfiler]:
|
||||
profiler: Final = PythonProfiler(source_root)
|
||||
def _installed_profiler(profiler: Callable[[FrameType, str, object], None], *, threads: bool) -> Generator[None]:
|
||||
if threads and sys.version_info >= (3, 12):
|
||||
tool_id: Final = next((slot for slot in (2, 3, 4, 0, 1, 5) if sys.monitoring.get_tool(slot) is None), None)
|
||||
if tool_id is None:
|
||||
raise RuntimeError("no sys.monitoring tool ID is available for Python trace collection")
|
||||
|
||||
def started(_code: CodeType, _offset: int) -> None:
|
||||
profiler(sys._getframe(1), "call", None)
|
||||
|
||||
sys.monitoring.use_tool_id(tool_id, "litellm-python-trace")
|
||||
try:
|
||||
sys.monitoring.register_callback(tool_id, sys.monitoring.events.PY_START, started)
|
||||
sys.monitoring.set_events(tool_id, sys.monitoring.events.PY_START)
|
||||
yield
|
||||
finally:
|
||||
sys.monitoring.set_events(tool_id, 0)
|
||||
sys.monitoring.register_callback(tool_id, sys.monitoring.events.PY_START, None)
|
||||
sys.monitoring.free_tool_id(tool_id)
|
||||
return
|
||||
if threads:
|
||||
warnings.warn(
|
||||
"Python <3.12 cannot trace existing worker threads; use Python 3.12+ for complete threaded traces",
|
||||
RuntimeWarning,
|
||||
stacklevel=3,
|
||||
)
|
||||
previous_thread: Final = threading.getprofile()
|
||||
if threads:
|
||||
threading.setprofile(profiler)
|
||||
previous: Final = sys.getprofile()
|
||||
sys.setprofile(profiler)
|
||||
try:
|
||||
yield profiler
|
||||
yield
|
||||
finally:
|
||||
sys.setprofile(previous)
|
||||
if threads:
|
||||
threading.setprofile(previous_thread)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def profile_python(source_root: Path, *, threads: bool = False) -> Generator[PythonProfiler]:
|
||||
profiler: Final = PythonProfiler(source_root)
|
||||
with _installed_profiler(profiler, threads=threads):
|
||||
yield profiler
|
||||
|
||||
|
||||
@contextmanager
|
||||
def profile_python_function_usage(
|
||||
source_root: Path,
|
||||
|
|
@ -160,14 +193,5 @@ def profile_python_function_usage(
|
|||
threads: bool = False,
|
||||
) -> Generator[PythonFunctionUsageProfiler]:
|
||||
profiler: Final = PythonFunctionUsageProfiler(source_root, functions)
|
||||
previous_thread: Final = threading.getprofile()
|
||||
if threads:
|
||||
threading.setprofile(profiler)
|
||||
previous: Final = sys.getprofile()
|
||||
sys.setprofile(profiler)
|
||||
try:
|
||||
with _installed_profiler(profiler, threads=threads):
|
||||
yield profiler
|
||||
finally:
|
||||
sys.setprofile(previous)
|
||||
if threads:
|
||||
threading.setprofile(previous_thread)
|
||||
|
|
|
|||
|
|
@ -76,7 +76,7 @@ def _span_for(engine: Engine, function: str, mappings: Sequence[TraceMapping]) -
|
|||
|
||||
|
||||
def pipeline_projection(
|
||||
engine: Engine, events: Sequence[FunctionTraceEvent], mappings: Sequence[TraceMapping]
|
||||
engine: Engine, events: Sequence[FunctionTraceEvent], mappings: Sequence[TraceMapping] | None = None
|
||||
) -> PipelineProjection:
|
||||
raw_parents: dict[int, int | None] = {}
|
||||
projected_ids: set[int] = set()
|
||||
|
|
@ -88,7 +88,7 @@ def pipeline_projection(
|
|||
if event.parent_id is not None and event.parent_id not in raw_parents:
|
||||
raise ValueError(f"trace event {event.id} references unknown or later parent {event.parent_id}")
|
||||
raw_parents[event.id] = event.parent_id
|
||||
span = _span_for(engine, event.function, mappings)
|
||||
span = event.function if mappings is None else _span_for(engine, event.function, mappings)
|
||||
if span is None:
|
||||
unmatched += 1
|
||||
continue
|
||||
|
|
@ -181,12 +181,7 @@ class TraceDiff:
|
|||
|
||||
@property
|
||||
def matches(self) -> bool:
|
||||
return (
|
||||
not self.python_only
|
||||
and not self.rust_only
|
||||
and not self.missing_mappings
|
||||
and self.shared_order_matches
|
||||
)
|
||||
return not self.python_only and not self.rust_only and not self.missing_mappings and self.shared_order_matches
|
||||
|
||||
|
||||
def _missing_mappings(
|
||||
|
|
@ -257,9 +252,7 @@ def trace_diff(
|
|||
rust_counts: Final = Counter(rust_spans)
|
||||
python_only_counts: Final = python_counts - rust_counts
|
||||
rust_only_counts: Final = rust_counts - python_counts
|
||||
python_only: Final = tuple(
|
||||
span for span, count in python_only_counts.items() for _ in range(count)
|
||||
)
|
||||
python_only: Final = tuple(span for span, count in python_only_counts.items() for _ in range(count))
|
||||
rust_only: Final = tuple(span for span, count in rust_only_counts.items() for _ in range(count))
|
||||
first_difference: Final = _first_difference(python, rust, mappings, contract)
|
||||
return TraceDiff(
|
||||
|
|
|
|||
|
|
@ -4,9 +4,10 @@ import asyncio
|
|||
import sys
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import wraps
|
||||
from pathlib import Path
|
||||
from types import FunctionType
|
||||
from types import FrameType, FunctionType
|
||||
from typing import Final, ParamSpec, TypeVar, cast
|
||||
|
||||
import pytest
|
||||
|
|
@ -41,11 +42,12 @@ def _events_named(profiler: PythonProfiler, name: str) -> tuple[FunctionTraceEve
|
|||
return tuple(event for event in profiler.events if event.function.endswith(name))
|
||||
|
||||
|
||||
def test_profiler_keeps_repeated_calls() -> None:
|
||||
@pytest.mark.parametrize("threads", (False, True))
|
||||
def test_profiler_keeps_repeated_calls(threads: bool) -> None:
|
||||
def called() -> None:
|
||||
return None
|
||||
|
||||
with profile_python(Path(__file__).parent) as profiler:
|
||||
with profile_python(Path(__file__).parent, threads=threads) as profiler:
|
||||
called()
|
||||
called()
|
||||
|
||||
|
|
@ -60,14 +62,15 @@ def test_profiler_qualifies_decorated_methods_by_class() -> None:
|
|||
assert _module_qualnames(__name__)[cast(FunctionType, Decorated.call.__wrapped__).__code__] == "Decorated.call"
|
||||
|
||||
|
||||
def test_profiler_records_real_frame_ancestry() -> None:
|
||||
@pytest.mark.parametrize("threads", (False, True))
|
||||
def test_profiler_records_real_frame_ancestry(threads: bool) -> None:
|
||||
def called() -> None:
|
||||
return None
|
||||
|
||||
def outer() -> None:
|
||||
called()
|
||||
|
||||
with profile_python(Path(__file__).parent) as profiler:
|
||||
with profile_python(Path(__file__).parent, threads=threads) as profiler:
|
||||
outer()
|
||||
|
||||
outer_event, called_event = (event for event in profiler.events if event.function.endswith(("outer", "called")))
|
||||
|
|
@ -84,18 +87,20 @@ def test_profiler_restores_previous_profiler_after_failure() -> None:
|
|||
assert sys.getprofile() is previous
|
||||
|
||||
|
||||
def test_profiler_does_not_count_coroutine_resumption_as_another_call() -> None:
|
||||
@pytest.mark.parametrize("threads", (False, True))
|
||||
def test_profiler_does_not_count_coroutine_resumption_as_another_call(threads: bool) -> None:
|
||||
async def suspended() -> None:
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
with profile_python(Path(__file__).parent) as profiler:
|
||||
with profile_python(Path(__file__).parent, threads=threads) as profiler:
|
||||
asyncio.run(suspended())
|
||||
|
||||
assert len(_events_named(profiler, "suspended")) == 1
|
||||
|
||||
|
||||
def test_profiler_preserves_parent_across_coroutine_suspension() -> None:
|
||||
@pytest.mark.parametrize("threads", (False, True))
|
||||
def test_profiler_preserves_parent_across_coroutine_suspension(threads: bool) -> None:
|
||||
def called() -> None:
|
||||
return None
|
||||
|
||||
|
|
@ -103,7 +108,7 @@ def test_profiler_preserves_parent_across_coroutine_suspension() -> None:
|
|||
await asyncio.sleep(0)
|
||||
called()
|
||||
|
||||
with profile_python(Path(__file__).parent) as profiler:
|
||||
with profile_python(Path(__file__).parent, threads=threads) as profiler:
|
||||
asyncio.run(suspended())
|
||||
|
||||
suspended_event: Final = _events_named(profiler, "suspended")[0]
|
||||
|
|
@ -124,6 +129,96 @@ def test_profiler_captures_worker_threads_when_enabled() -> None:
|
|||
assert called_event.parent_id is None
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="existing worker capture requires sys.monitoring")
|
||||
@pytest.mark.parametrize("prewarm", (False, True))
|
||||
def test_profiler_captures_reused_workers_without_leaking_between_sessions(prewarm: bool) -> None:
|
||||
def called() -> None:
|
||||
return None
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
if prewarm:
|
||||
executor.submit(called).result(timeout=5)
|
||||
with profile_python(Path(__file__).parent, threads=True) as first:
|
||||
executor.submit(called).result(timeout=5)
|
||||
executor.submit(called).result(timeout=5)
|
||||
with profile_python(Path(__file__).parent, threads=True) as second:
|
||||
executor.submit(called).result(timeout=5)
|
||||
executor.submit(called).result(timeout=5)
|
||||
|
||||
assert len(_events_named(first, "called")) == 1
|
||||
assert len(_events_named(second, "called")) == 1
|
||||
|
||||
|
||||
def test_profiler_restores_main_and_worker_hooks_after_failure() -> None:
|
||||
previous: Final = sys.getprofile()
|
||||
previous_thread: Final = threading.getprofile()
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
worker_previous: Final = executor.submit(sys.getprofile).result(timeout=5)
|
||||
with pytest.raises(RuntimeError, match="stop"):
|
||||
with profile_python(Path(__file__).parent, threads=True):
|
||||
raise RuntimeError("stop")
|
||||
assert executor.submit(sys.getprofile).result(timeout=5) is worker_previous
|
||||
|
||||
assert sys.getprofile() is previous
|
||||
assert threading.getprofile() is previous_thread
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="existing worker capture requires sys.monitoring")
|
||||
def test_function_usage_profiler_captures_reused_workers() -> None:
|
||||
def selected() -> None:
|
||||
return None
|
||||
|
||||
function: Final = f"{Path(__file__).name}:{selected.__code__.co_firstlineno} {selected.__qualname__}"
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
executor.submit(selected).result(timeout=5)
|
||||
with profile_python_function_usage(Path(__file__).parent, frozenset((function,)), threads=True) as profiler:
|
||||
executor.submit(selected).result(timeout=5)
|
||||
|
||||
assert profiler.called == {function}
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="independent thread hooks require sys.monitoring")
|
||||
def test_threaded_profiler_preserves_custom_worker_hook_and_releases_monitoring_slot() -> None:
|
||||
def worker_hook(_frame: FrameType, _event: str, _arg: object) -> None:
|
||||
return None
|
||||
|
||||
def fail_with_profile(executor: ThreadPoolExecutor) -> None:
|
||||
with profile_python(Path(__file__).parent, threads=True):
|
||||
assert executor.submit(sys.getprofile).result(timeout=5) is worker_hook
|
||||
raise RuntimeError("stop")
|
||||
|
||||
tools_before: Final = tuple(sys.monitoring.get_tool(slot) for slot in range(6))
|
||||
with ThreadPoolExecutor(max_workers=1, initializer=lambda: sys.setprofile(worker_hook)) as executor:
|
||||
assert executor.submit(sys.getprofile).result(timeout=5) is worker_hook
|
||||
with pytest.raises(RuntimeError, match="stop"):
|
||||
fail_with_profile(executor)
|
||||
assert executor.submit(sys.getprofile).result(timeout=5) is worker_hook
|
||||
|
||||
assert tuple(sys.monitoring.get_tool(slot) for slot in range(6)) == tools_before
|
||||
|
||||
|
||||
@pytest.mark.skipif(sys.version_info < (3, 12), reason="existing worker capture requires sys.monitoring")
|
||||
def test_threaded_profiler_keeps_concurrent_event_ids_and_parent_links() -> None:
|
||||
def child() -> None:
|
||||
return None
|
||||
|
||||
def parent() -> None:
|
||||
child()
|
||||
|
||||
with ThreadPoolExecutor(max_workers=4) as executor:
|
||||
with profile_python(Path(__file__).parent, threads=True) as profiler:
|
||||
futures: Final = tuple(executor.submit(parent) for _ in range(200))
|
||||
for future in futures:
|
||||
future.result(timeout=5)
|
||||
|
||||
parent_ids: Final = frozenset(event.id for event in _events_named(profiler, "parent"))
|
||||
children: Final = _events_named(profiler, "child")
|
||||
assert len(parent_ids) == len(children) == 200
|
||||
assert frozenset(event.parent_id for event in children) == parent_ids
|
||||
assert tuple(event.id for event in profiler.events) == tuple(range(len(profiler.events)))
|
||||
|
||||
|
||||
def test_function_usage_profiler_records_only_selected_functions() -> None:
|
||||
def selected() -> None:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -40,6 +40,23 @@ def test_python_projection_collapses_unmapped_parents_and_counts_noise() -> None
|
|||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("engine", ("python", "rust"))
|
||||
def test_projection_without_mappings_keeps_every_call_and_parent(engine: Engine) -> None:
|
||||
events: Final = (
|
||||
event(0, "module.py:1 entry"),
|
||||
event(1, "module.py:2 internal_helper", 0),
|
||||
event(2, "module.py:3 nested", 1),
|
||||
event(3, "module.py:2 internal_helper", 0),
|
||||
)
|
||||
|
||||
projection: Final = pipeline_projection(engine, events)
|
||||
|
||||
assert projection.unmatched == 0
|
||||
assert tuple((step.id, step.parent_id, step.span, step.raw) for step in projection.steps) == tuple(
|
||||
(item.id, item.parent_id, item.function, item.raw) for item in events
|
||||
)
|
||||
|
||||
|
||||
def test_rust_projection_keeps_unknown_spans() -> None:
|
||||
projection: Final = pipeline_projection("rust", (event(0, "route"), event(1, "new_span", 0)), MAPPINGS)
|
||||
assert [(step.span, step.parent_id) for step in projection.steps] == [("route", None), ("new_span", 0)]
|
||||
|
|
@ -146,9 +163,7 @@ def test_trace_diff_allows_reordered_concurrent_children() -> None:
|
|||
def test_trace_diff_prunes_declared_engine_only_nodes_but_requires_them() -> None:
|
||||
mappings: Final = (MAPPINGS[0], mapping(rust_span="rust_prepare"))
|
||||
python: Final = pipeline_projection("python", (event(0, "module.py:1 entry"),), mappings).steps
|
||||
rust: Final = pipeline_projection(
|
||||
"rust", (event(0, "route"), event(1, "rust_prepare", 0)), mappings
|
||||
).steps
|
||||
rust: Final = pipeline_projection("rust", (event(0, "route"), event(1, "rust_prepare", 0)), mappings).steps
|
||||
assert trace_diff(python, rust, mappings).matches
|
||||
assert trace_diff(python, rust[:1], mappings).missing_mappings == ("rust_prepare",)
|
||||
|
||||
|
|
|
|||
|
|
@ -264,9 +264,7 @@ def _formatting_strategy() -> SearchStrategy[ReductoFormatting]:
|
|||
),
|
||||
st.sampled_from((False, True)).map(lambda value: {"add_page_markers": value}),
|
||||
st.sampled_from((False, True)).map(lambda value: {"merge_tables": value}),
|
||||
st.sampled_from(REDUCTO_FORMATTING_INCLUDE_GROUPS)
|
||||
.map(list)
|
||||
.map(lambda value: {"include": value}),
|
||||
st.sampled_from(REDUCTO_FORMATTING_INCLUDE_GROUPS).map(list).map(lambda value: {"include": value}),
|
||||
)
|
||||
return values.map(ReductoFormatting.model_validate)
|
||||
|
||||
|
|
|
|||
|
|
@ -124,10 +124,7 @@ class RecordingCallback(CustomLogger):
|
|||
if isinstance(value, Mapping):
|
||||
if any(not isinstance(map_key, str) for map_key in value):
|
||||
raise TypeError("callback kwarg mappings must use string keys")
|
||||
return {
|
||||
map_key: self._normalized_kwargs(map_value, map_key)
|
||||
for map_key, map_value in value.items()
|
||||
}
|
||||
return {map_key: self._normalized_kwargs(map_value, map_key) for map_key, map_value in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [self._normalized_kwargs(item) for item in value]
|
||||
raise TypeError(f"unsupported callback kwarg type: {type(value)}")
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
Maps Python profiler frames onto feature-gated Rust span names via an explicit per-case mapping (Rust span name is the identity) and compares steps, order, and nesting of both live traces against a replayed provider response.
|
||||
Prints every collected Python call under litellm/ and every feature-gated Rust span from live traces against replayed HTTP responses. The two traces are independent and are not compared. API-key and Vertex credentials scenarios exercise separate authentication paths; credentials scenarios replay the token exchange locally.
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from ...shared.reporting.strategy import (
|
|||
ModuleCaseSpec,
|
||||
NotImplementedCaseSpec,
|
||||
RunnerArgumentDefinition,
|
||||
RunnerOptionDefinition,
|
||||
StrategyDefinition,
|
||||
)
|
||||
from .reporting import render_trace_results
|
||||
|
|
@ -26,13 +27,17 @@ CASES: Final[tuple[CaseDefinition, ...]] = (
|
|||
ModuleCaseSpec(
|
||||
coverage=Coverage.PARTIAL,
|
||||
module="tests.rust-python-harness.strategies.trace_parity.sdk.messages.case",
|
||||
note="Async only until anthropic_messages_handler supports sync calls.",
|
||||
note="Success paths are async; sync tracing captures the currently unsupported behavior.",
|
||||
),
|
||||
surface="sdk",
|
||||
),
|
||||
CaseDefinition(
|
||||
"responses",
|
||||
NotImplementedCaseSpec(reason="No Responses trace-parity case is registered."),
|
||||
ModuleCaseSpec(
|
||||
coverage=Coverage.PARTIAL,
|
||||
module="tests.rust-python-harness.strategies.trace_parity.sdk.responses.case",
|
||||
note="Core create paths: native, streaming, provider error, Azure override, and chat bridge.",
|
||||
),
|
||||
surface="sdk",
|
||||
),
|
||||
CaseDefinition(
|
||||
|
|
@ -70,13 +75,17 @@ CASES: Final[tuple[CaseDefinition, ...]] = (
|
|||
ModuleCaseSpec(
|
||||
coverage=Coverage.PARTIAL,
|
||||
module="tests.rust-python-harness.strategies.trace_parity.gateway.messages.case",
|
||||
note="Non-streaming success paths only.",
|
||||
note="Anthropic/Azure provider routes plus a fully consumed downstream streaming path.",
|
||||
),
|
||||
surface="gateway",
|
||||
),
|
||||
CaseDefinition(
|
||||
"responses",
|
||||
NotImplementedCaseSpec(reason="No gateway Responses trace-parity case is registered."),
|
||||
ModuleCaseSpec(
|
||||
coverage=Coverage.PARTIAL,
|
||||
module="tests.rust-python-harness.strategies.trace_parity.gateway.responses.case",
|
||||
note="Native OpenAI non-streaming and fully consumed downstream streaming paths.",
|
||||
),
|
||||
surface="gateway",
|
||||
),
|
||||
CaseDefinition(
|
||||
|
|
@ -86,7 +95,11 @@ CASES: Final[tuple[CaseDefinition, ...]] = (
|
|||
),
|
||||
CaseDefinition(
|
||||
"chat_completions",
|
||||
NotImplementedCaseSpec(reason="No gateway chat trace-parity case is registered."),
|
||||
ModuleCaseSpec(
|
||||
coverage=Coverage.PARTIAL,
|
||||
module="tests.rust-python-harness.strategies.trace_parity.gateway.chat_completions.case",
|
||||
note="Anthropic non-streaming and fully consumed downstream streaming paths.",
|
||||
),
|
||||
surface="gateway",
|
||||
),
|
||||
CaseDefinition(
|
||||
|
|
@ -99,8 +112,8 @@ CASES: Final[tuple[CaseDefinition, ...]] = (
|
|||
STRATEGY: Final = StrategyDefinition(
|
||||
id="trace_parity",
|
||||
order=20,
|
||||
label="Trace parity",
|
||||
description="Compare pipeline steps, order, and nesting between Python profiler frames and Rust spans via an explicit mapping.",
|
||||
label="Traces",
|
||||
description="Print Python profiler frames and Rust spans for representative pipeline scenarios.",
|
||||
directory=Path(__file__).parent,
|
||||
runnable_spec=ModuleCaseSpec,
|
||||
cases=CASES,
|
||||
|
|
@ -112,4 +125,11 @@ STRATEGY: Final = StrategyDefinition(
|
|||
metavar="NAME",
|
||||
help="run only this named trace scenario; repeat to select more than one",
|
||||
),
|
||||
runner_options=(
|
||||
RunnerOptionDefinition(
|
||||
option="--engine",
|
||||
choices=("python", "rust"),
|
||||
help="show only this engine's trace; omit to print both engines",
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
177
tests/rust-python-harness/strategies/trace_parity/fixtures.py
Normal file
177
tests/rust-python-harness/strategies/trace_parity/fixtures.py
Normal file
|
|
@ -0,0 +1,177 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
import json
|
||||
import struct
|
||||
from collections.abc import Iterable, Mapping
|
||||
from typing import Final
|
||||
|
||||
from ...shared.parity.recorded_http import (
|
||||
HttpHeader,
|
||||
RecordedHttpResponse,
|
||||
RecordedHttpStreamResponse,
|
||||
RecordedStreamChunk,
|
||||
)
|
||||
|
||||
JSON_HEADERS: Final = (HttpHeader(name="content-type", value="application/json"),)
|
||||
SSE_HEADERS: Final = (HttpHeader(name="content-type", value="text/event-stream"),)
|
||||
AWS_EVENT_STREAM_HEADERS: Final = (HttpHeader(name="content-type", value="application/vnd.amazon.eventstream"),)
|
||||
|
||||
|
||||
def json_response(body: Mapping[str, object] | bytes, *, status: int = 200) -> RecordedHttpResponse:
|
||||
encoded: Final = body if isinstance(body, bytes) else json.dumps(body).encode()
|
||||
return RecordedHttpResponse.from_bytes(status, JSON_HEADERS, encoded)
|
||||
|
||||
|
||||
def sse_event(event: str, payload: Mapping[str, object]) -> bytes:
|
||||
return f"event: {event}\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n".encode()
|
||||
|
||||
|
||||
def sse_response(events: Iterable[tuple[str, Mapping[str, object]]]) -> RecordedHttpStreamResponse:
|
||||
return RecordedHttpStreamResponse(
|
||||
kind="http_stream",
|
||||
status_code=200,
|
||||
headers=SSE_HEADERS,
|
||||
chunks=tuple(RecordedStreamChunk.from_bytes(sse_event(event, payload)) for event, payload in events),
|
||||
)
|
||||
|
||||
|
||||
def _aws_string_header(name: str, value: str) -> bytes:
|
||||
name_bytes: Final = name.encode()
|
||||
value_bytes: Final = value.encode()
|
||||
return (
|
||||
struct.pack("!B", len(name_bytes))
|
||||
+ name_bytes
|
||||
+ struct.pack("!B", 7)
|
||||
+ struct.pack("!H", len(value_bytes))
|
||||
+ value_bytes
|
||||
)
|
||||
|
||||
|
||||
def aws_event_stream_frame(payload: Mapping[str, object]) -> bytes:
|
||||
event_payload: Final = json.dumps(
|
||||
{"bytes": base64.b64encode(json.dumps(payload, separators=(",", ":")).encode()).decode()},
|
||||
separators=(",", ":"),
|
||||
).encode()
|
||||
headers: Final = (
|
||||
_aws_string_header(":event-type", "chunk")
|
||||
+ _aws_string_header(":content-type", "application/json")
|
||||
+ _aws_string_header(":message-type", "event")
|
||||
)
|
||||
total_length: Final = 12 + len(headers) + len(event_payload) + 4
|
||||
prelude: Final = struct.pack("!II", total_length, len(headers))
|
||||
prelude_crc: Final = binascii.crc32(prelude) & 0xFFFFFFFF
|
||||
prelude_crc_bytes: Final = struct.pack("!I", prelude_crc)
|
||||
message_crc: Final = binascii.crc32(prelude_crc_bytes + headers + event_payload, prelude_crc) & 0xFFFFFFFF
|
||||
return prelude + prelude_crc_bytes + headers + event_payload + struct.pack("!I", message_crc)
|
||||
|
||||
|
||||
def aws_event_stream_response(
|
||||
events: Iterable[Mapping[str, object]], *, corrupt_last_frame: bool = False
|
||||
) -> RecordedHttpStreamResponse:
|
||||
frames: Final = tuple(aws_event_stream_frame(event) for event in events)
|
||||
body: Final = (
|
||||
b"".join((*frames[:-1], frames[-1][:-1] + bytes((frames[-1][-1] ^ 0xFF,))))
|
||||
if corrupt_last_frame
|
||||
else b"".join(frames)
|
||||
)
|
||||
return RecordedHttpStreamResponse(
|
||||
kind="http_stream",
|
||||
status_code=200,
|
||||
headers=AWS_EVENT_STREAM_HEADERS,
|
||||
chunks=(RecordedStreamChunk.from_bytes(body),),
|
||||
)
|
||||
|
||||
|
||||
def anthropic_response_body(*, model: str = "claude-sonnet-5") -> dict[str, object]:
|
||||
return {
|
||||
"id": "msg_trace",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": model,
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 3},
|
||||
}
|
||||
|
||||
|
||||
def anthropic_stream_events(*, model: str = "claude-sonnet-5") -> tuple[tuple[str, Mapping[str, object]], ...]:
|
||||
return (
|
||||
(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_trace",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": model,
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 0},
|
||||
},
|
||||
},
|
||||
),
|
||||
(
|
||||
"content_block_start",
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
),
|
||||
(
|
||||
"content_block_delta",
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello"}},
|
||||
),
|
||||
("content_block_stop", {"type": "content_block_stop", "index": 0}),
|
||||
(
|
||||
"message_delta",
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
|
||||
"usage": {"output_tokens": 1},
|
||||
},
|
||||
),
|
||||
("message_stop", {"type": "message_stop"}),
|
||||
)
|
||||
|
||||
|
||||
def responses_body(*, model: str = "gpt-5", status: str = "completed") -> dict[str, object]:
|
||||
return {
|
||||
"id": "resp_trace",
|
||||
"object": "response",
|
||||
"created_at": 1_750_000_000,
|
||||
"status": status,
|
||||
"model": model,
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_trace",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "hello", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 2, "output_tokens": 3, "total_tokens": 5},
|
||||
}
|
||||
|
||||
|
||||
def responses_stream_events(*, model: str = "gpt-5") -> tuple[tuple[str, Mapping[str, object]], ...]:
|
||||
response: Final = responses_body(model=model)
|
||||
return (
|
||||
(
|
||||
"response.created",
|
||||
{"type": "response.created", "response": {**response, "status": "in_progress", "output": []}},
|
||||
),
|
||||
(
|
||||
"response.output_text.delta",
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": "msg_trace",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": "hello",
|
||||
},
|
||||
),
|
||||
("response.completed", {"type": "response.completed", "response": response}),
|
||||
)
|
||||
|
|
@ -0,0 +1,63 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
from .....shared.tracing.steps import Engine, mapping
|
||||
from ...fixtures import anthropic_response_body, anthropic_stream_events, json_response, sse_response
|
||||
from ...models import GatewayRouteSpec, RouteFixture, TraceScenario, TraceSuite
|
||||
|
||||
MAPPINGS: Final = (
|
||||
mapping(span="python_chat_gateway_route", python_frame=r"proxy_server\.py:\d+ chat_completion$"),
|
||||
mapping(span="python_gateway_service", python_frame=r"ProxyBaseLLMRequestProcessing\.base_process_llm_request$"),
|
||||
mapping(span="python_chat_entrypoint", python_frame=r"main\.py:\d+ a?completion$"),
|
||||
mapping(span="python_provider_config", python_frame=r"ProviderConfigManager\.get_provider_chat_config$"),
|
||||
mapping(rust_span="validate_environment", python_frame=r"(?<!_)validate_environment$"),
|
||||
mapping(rust_span="transform_request", python_frame=r"AnthropicConfig\.transform_request$"),
|
||||
mapping(span="python_logging_pre_call", python_frame=r"Logging\.pre_call$"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(rust_span="transform_response", python_frame=r"AnthropicConfig\.transform_response$"),
|
||||
mapping(span="python_success_callback", python_frame=r"Logging\.async_success_handler$|Logging\.success_handler$"),
|
||||
)
|
||||
|
||||
STREAM_MAPPINGS: Final = (
|
||||
mapping(span="python_stream_wrapper", python_frame=r"CustomStreamWrapper\.__init__$"),
|
||||
mapping(span="python_stream_next", python_frame=r"CustomStreamWrapper\.__anext__$"),
|
||||
mapping(span="python_stream_chunk", python_frame=r"CustomStreamWrapper\.chunk_creator$"),
|
||||
mapping(span="python_downstream_stream", python_frame=r"DataGenerator\.__anext__$|async_data_generator$"),
|
||||
)
|
||||
|
||||
|
||||
def _fixture(_engine: Engine, _base_url: str) -> RouteFixture:
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
"model_alias": "trace-model",
|
||||
"provider_model": "anthropic/claude-sonnet-5",
|
||||
"body": {
|
||||
"model": "trace-model",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 16,
|
||||
},
|
||||
},
|
||||
provider_responses=(json_response(anthropic_response_body()),),
|
||||
)
|
||||
|
||||
|
||||
def _stream_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
fixture: Final = _fixture(engine, base_url)
|
||||
return fixture.with_body(stream=True).derive(
|
||||
provider_responses=(sse_response(anthropic_stream_events()),),
|
||||
)
|
||||
|
||||
|
||||
TRACE_SUITE: Final = TraceSuite(
|
||||
route=GatewayRouteSpec("chat_completions", rust_supported=False),
|
||||
scenarios=(
|
||||
TraceScenario(name="async-anthropic", fixture=_fixture, mappings=MAPPINGS, asynchronous=True),
|
||||
TraceScenario(
|
||||
name="async-anthropic-downstream-stream",
|
||||
fixture=_stream_fixture,
|
||||
mappings=(*MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
@ -13,8 +13,8 @@ from ....shared.parity.replay import replay_server
|
|||
from ....shared.tracing.native import TraceResponsePayload, native_trace_events
|
||||
from ....shared.tracing.profiler import FunctionTraceEvent, profile_python
|
||||
from ....shared.tracing.steps import Engine, PipelineProjection, pipeline_projection
|
||||
from ..models import GatewayRouteSpec, RouteFixture, TraceExecutionFailure, TraceMode, TraceScenario
|
||||
from ..reporting import TraceComparisonArtifact
|
||||
from ..models import GatewayRouteSpec, RouteFixture, TraceEngine, TraceExecutionFailure, TraceScenario
|
||||
from ..reporting import TraceArtifact
|
||||
|
||||
|
||||
class _GatewayResponsePayload(BaseModel):
|
||||
|
|
@ -28,7 +28,14 @@ class _GatewayClient(Protocol):
|
|||
def post(self, url: str, *, json: object, headers: dict[str, str]) -> httpx.Response: ...
|
||||
|
||||
|
||||
def _collect_python(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]:
|
||||
_ROUTE_PATHS: Final = {
|
||||
"messages": "/v1/messages",
|
||||
"chat_completions": "/v1/chat/completions",
|
||||
"responses": "/v1/responses",
|
||||
}
|
||||
|
||||
|
||||
def _collect_python(fixture: RouteFixture, route: GatewayRouteSpec) -> tuple[FunctionTraceEvent, ...]:
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm
|
||||
|
|
@ -61,7 +68,7 @@ def _collect_python(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]:
|
|||
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
|
||||
client: Final = cast(_GatewayClient, TestClient(proxy_server.app))
|
||||
response: Final = client.post(
|
||||
"/v1/messages",
|
||||
_ROUTE_PATHS[route.route],
|
||||
json=fixture.kwargs["body"],
|
||||
headers={"authorization": "Bearer trace-key"},
|
||||
)
|
||||
|
|
@ -76,9 +83,10 @@ def _collect_python(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]:
|
|||
proxy_server.app.dependency_overrides[user_api_key_auth] = old_override
|
||||
|
||||
|
||||
def _collect_rust(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]:
|
||||
def _collect_rust(fixture: RouteFixture, route: GatewayRouteSpec) -> tuple[FunctionTraceEvent, ...]:
|
||||
payload: Final = json.dumps(
|
||||
{
|
||||
"path": _ROUTE_PATHS[route.route],
|
||||
"model_alias": fixture.kwargs["model_alias"],
|
||||
"provider_model": fixture.kwargs["provider_model"],
|
||||
"api_base": fixture.kwargs["api_base"],
|
||||
|
|
@ -130,7 +138,9 @@ def _gateway_trace_binary() -> Path:
|
|||
return rust_root / "target" / "debug" / "trace-parity-gateway"
|
||||
|
||||
|
||||
def _collect(scenario: TraceScenario, engine: Engine) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
|
||||
def _collect(
|
||||
route: GatewayRouteSpec, scenario: TraceScenario, engine: Engine
|
||||
) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
|
||||
try:
|
||||
with replay_server() as provider:
|
||||
base_fixture: Final = scenario.fixture(engine, provider.url)
|
||||
|
|
@ -140,7 +150,7 @@ def _collect(scenario: TraceScenario, engine: Engine) -> tuple[FunctionTraceEven
|
|||
)
|
||||
for response in fixture.provider_responses:
|
||||
provider.enqueue_response(response)
|
||||
events: Final = _collect_python(fixture) if engine == "python" else _collect_rust(fixture)
|
||||
events: Final = _collect_python(fixture, route) if engine == "python" else _collect_rust(fixture, route)
|
||||
provider.take_requests(len(fixture.provider_responses))
|
||||
return events
|
||||
except Exception as error:
|
||||
|
|
@ -150,40 +160,38 @@ def _collect(scenario: TraceScenario, engine: Engine) -> tuple[FunctionTraceEven
|
|||
def _projections(
|
||||
python_events: tuple[FunctionTraceEvent, ...],
|
||||
rust_events: tuple[FunctionTraceEvent, ...],
|
||||
scenario: TraceScenario,
|
||||
mode: TraceMode,
|
||||
) -> tuple[PipelineProjection, PipelineProjection, str | None]:
|
||||
mappings: Final = scenario.mappings_for(mode)
|
||||
try:
|
||||
return (
|
||||
pipeline_projection("python", python_events, mappings),
|
||||
pipeline_projection("rust", rust_events, mappings),
|
||||
pipeline_projection("python", python_events),
|
||||
pipeline_projection("rust", rust_events),
|
||||
None,
|
||||
)
|
||||
except ValueError as error:
|
||||
return PipelineProjection(), PipelineProjection(), f"harness: {error}"
|
||||
|
||||
|
||||
def execute_gateway_trace(route: GatewayRouteSpec, scenario: TraceScenario, mode: TraceMode) -> TraceComparisonArtifact:
|
||||
mappings: Final = scenario.mappings_for(mode)
|
||||
python_trace: Final = _collect(scenario, "python")
|
||||
rust_trace: Final = _collect(scenario, "rust")
|
||||
def execute_gateway_trace(
|
||||
route: GatewayRouteSpec,
|
||||
scenario: TraceScenario,
|
||||
engine: TraceEngine = "both",
|
||||
) -> TraceArtifact:
|
||||
effective_engine: Final[TraceEngine] = "python" if engine == "both" and not route.rust_supported else engine
|
||||
python_trace: Final = _collect(route, scenario, "python") if effective_engine != "rust" else ()
|
||||
rust_trace: Final = _collect(route, scenario, "rust") if effective_engine != "python" else ()
|
||||
collection_python_error: Final = None if isinstance(python_trace, tuple) else f"python: {python_trace.message}"
|
||||
rust_error: Final = None if isinstance(rust_trace, tuple) else f"rust: {rust_trace.message}"
|
||||
python_events: Final = python_trace if isinstance(python_trace, tuple) else ()
|
||||
rust_events: Final = rust_trace if isinstance(rust_trace, tuple) else ()
|
||||
python, rust, projection_error = _projections(python_events, rust_events, scenario, mode)
|
||||
python, rust, projection_error = _projections(python_events, rust_events)
|
||||
python_error: Final = projection_error or collection_python_error
|
||||
return TraceComparisonArtifact.from_traces(
|
||||
return TraceArtifact.from_traces(
|
||||
engine=effective_engine,
|
||||
surface="gateway",
|
||||
sdk_function=route.route,
|
||||
scenario=scenario.name,
|
||||
mode=mode,
|
||||
mappings=mappings,
|
||||
contract=scenario.contract,
|
||||
python=python.steps,
|
||||
rust=rust.steps,
|
||||
python_unmatched=python.unmatched,
|
||||
python_error=python_error,
|
||||
rust_error=rust_error,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
from .....shared.parity.recorded_http import HttpHeader, RecordedHttpResponse
|
||||
from .....shared.tracing.steps import Engine, mapping
|
||||
from ...fixtures import anthropic_response_body, anthropic_stream_events, json_response, sse_response
|
||||
from ...models import GatewayRouteSpec, RouteFixture, TraceScenario, TraceSuite
|
||||
|
||||
|
||||
|
|
@ -49,24 +48,7 @@ def _fixture(_engine: Engine, provider: str) -> RouteFixture:
|
|||
"max_tokens": 16,
|
||||
},
|
||||
},
|
||||
provider_responses=(
|
||||
RecordedHttpResponse.from_bytes(
|
||||
200,
|
||||
(HttpHeader(name="content-type", value="application/json"),),
|
||||
json.dumps(
|
||||
{
|
||||
"id": "msg_trace",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-5",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 3},
|
||||
}
|
||||
).encode(),
|
||||
),
|
||||
),
|
||||
provider_responses=(json_response(anthropic_response_body()),),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -78,6 +60,13 @@ def _azure_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
|||
return _fixture(engine, "azure_ai")
|
||||
|
||||
|
||||
def _stream_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _anthropic_fixture(engine, _base_url)
|
||||
return fixture.with_body(stream=True).derive(
|
||||
provider_responses=(sse_response(anthropic_stream_events()),),
|
||||
)
|
||||
|
||||
|
||||
ANTHROPIC_MAPPINGS: Final = (
|
||||
*GATEWAY_MAPPINGS,
|
||||
mapping(
|
||||
|
|
@ -100,7 +89,22 @@ AZURE_MAPPINGS: Final = (
|
|||
TRACE_SUITE: Final = TraceSuite(
|
||||
route=GatewayRouteSpec("messages"),
|
||||
scenarios=(
|
||||
TraceScenario(name="anthropic", fixture=_anthropic_fixture, mappings=ANTHROPIC_MAPPINGS, modes=("async",)),
|
||||
TraceScenario(name="azure-ai", fixture=_azure_fixture, mappings=AZURE_MAPPINGS, modes=("async",)),
|
||||
TraceScenario(
|
||||
name="async-anthropic", fixture=_anthropic_fixture, mappings=ANTHROPIC_MAPPINGS, asynchronous=True
|
||||
),
|
||||
TraceScenario(name="async-azure-ai", fixture=_azure_fixture, mappings=AZURE_MAPPINGS, asynchronous=True),
|
||||
TraceScenario(
|
||||
name="async-anthropic-downstream-stream",
|
||||
fixture=_stream_fixture,
|
||||
mappings=(
|
||||
*ANTHROPIC_MAPPINGS,
|
||||
mapping(span="python_upstream_stream", python_frame=r"AnthropicMessagesStreamingResponse\.__anext__$"),
|
||||
mapping(
|
||||
span="python_downstream_stream", python_frame=r"DataGenerator\.__anext__$|async_data_generator$"
|
||||
),
|
||||
mapping(span="python_stream_callback", python_frame=r"Logging\.async_success_handler$"),
|
||||
),
|
||||
asynchronous=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,62 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
from .....shared.tracing.steps import Engine, mapping
|
||||
from ...fixtures import json_response, responses_body, responses_stream_events, sse_response
|
||||
from ...models import GatewayRouteSpec, RouteFixture, TraceScenario, TraceSuite
|
||||
|
||||
MAPPINGS: Final = (
|
||||
mapping(
|
||||
span="python_responses_gateway_route", python_frame=r"response_api_endpoints/endpoints\.py:\d+ responses_api$"
|
||||
),
|
||||
mapping(span="python_gateway_service", python_frame=r"ProxyBaseLLMRequestProcessing\.base_process_llm_request$"),
|
||||
mapping(span="python_responses", python_frame=r"responses/main\.py:\d+ a?responses$"),
|
||||
mapping(span="python_provider_config", python_frame=r"ProviderConfigManager\.get_provider_responses_api_config$"),
|
||||
mapping(rust_span="validate_environment", python_frame=r"OpenAIResponsesAPIConfig\.validate_environment$"),
|
||||
mapping(rust_span="complete_url", python_frame=r"OpenAIResponsesAPIConfig\.get_complete_url$"),
|
||||
mapping(rust_span="transform_request", python_frame=r"OpenAIResponsesAPIConfig\.transform_responses_api_request$"),
|
||||
mapping(span="python_logging_pre_call", python_frame=r"Logging\.pre_call$"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(rust_span="transform_response", python_frame=r"OpenAIResponsesAPIConfig\.transform_response_api_response$"),
|
||||
mapping(span="python_success_callback", python_frame=r"Logging\.async_success_handler$|Logging\.success_handler$"),
|
||||
)
|
||||
|
||||
STREAM_MAPPINGS: Final = (
|
||||
mapping(span="python_stream_iterator", python_frame=r"ResponsesAPIStreamingIterator\.__init__$"),
|
||||
mapping(span="python_stream_next", python_frame=r"ResponsesAPIStreamingIterator\.__anext__$"),
|
||||
mapping(span="python_stream_transform", python_frame=r"OpenAIResponsesAPIConfig\.transform_streaming_response$"),
|
||||
mapping(span="python_downstream_stream", python_frame=r"DataGenerator\.__anext__$|async_data_generator$"),
|
||||
)
|
||||
|
||||
|
||||
def _fixture(_engine: Engine, _base_url: str) -> RouteFixture:
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
"model_alias": "trace-model",
|
||||
"provider_model": "openai/gpt-5",
|
||||
"body": {"model": "trace-model", "input": "hello"},
|
||||
},
|
||||
provider_responses=(json_response(responses_body()),),
|
||||
)
|
||||
|
||||
|
||||
def _stream_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
fixture: Final = _fixture(engine, base_url)
|
||||
return fixture.with_body(stream=True).derive(
|
||||
provider_responses=(sse_response(responses_stream_events()),),
|
||||
)
|
||||
|
||||
|
||||
TRACE_SUITE: Final = TraceSuite(
|
||||
route=GatewayRouteSpec("responses", rust_supported=False),
|
||||
scenarios=(
|
||||
TraceScenario(name="async-openai", fixture=_fixture, mappings=MAPPINGS, asynchronous=True),
|
||||
TraceScenario(
|
||||
name="async-openai-downstream-stream",
|
||||
fixture=_stream_fixture,
|
||||
mappings=(*MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
@ -1,35 +1,61 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Literal, TypeAlias
|
||||
from typing import Final, Literal, TypeAlias, cast
|
||||
|
||||
from ...shared.parity.recorded_http import RecordedHttpResponse
|
||||
from ...shared.parity.recorded_http import RecordedResponse
|
||||
from ...shared.reporting.models import SdkFunction
|
||||
from ...shared.tracing.steps import Engine, TraceContract, TraceMapping
|
||||
from ...shared.tracing.steps import Engine, TraceMapping
|
||||
|
||||
TraceMode = Literal["sync", "async"]
|
||||
TraceEngine = Literal["python", "rust", "both"]
|
||||
TraceFailureSource = Literal["python", "rust", "harness"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouteFixture:
|
||||
kwargs: dict[str, object]
|
||||
provider_responses: tuple[RecordedHttpResponse, ...]
|
||||
provider_responses: tuple[RecordedResponse, ...]
|
||||
expected_failure: bool = False
|
||||
consume_stream: bool = False
|
||||
environment: tuple[tuple[str, str], ...] = ()
|
||||
|
||||
def derive(
|
||||
self,
|
||||
*,
|
||||
kwargs: Mapping[str, object] | None = None,
|
||||
provider_responses: tuple[RecordedResponse, ...] | None = None,
|
||||
expected_failure: bool | None = None,
|
||||
consume_stream: bool | None = None,
|
||||
) -> RouteFixture:
|
||||
return RouteFixture(
|
||||
kwargs={**self.kwargs, **(kwargs or {})},
|
||||
provider_responses=self.provider_responses if provider_responses is None else provider_responses,
|
||||
expected_failure=self.expected_failure if expected_failure is None else expected_failure,
|
||||
consume_stream=self.consume_stream if consume_stream is None else consume_stream,
|
||||
environment=self.environment,
|
||||
)
|
||||
|
||||
def with_body(self, **updates: object) -> RouteFixture:
|
||||
raw_body: Final = self.kwargs.get("body")
|
||||
if not isinstance(raw_body, dict):
|
||||
raise ValueError("route fixture does not contain an object body")
|
||||
body: Final = cast(dict[str, object], raw_body)
|
||||
return self.derive(kwargs={"body": {**body, **updates}})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouteSpec:
|
||||
route: SdkFunction
|
||||
python_entrypoints: tuple[str, str]
|
||||
rust_entrypoints: tuple[str, str]
|
||||
rust_entrypoints: tuple[str, str] | None
|
||||
fixture: Callable[[Engine, str], RouteFixture]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GatewayRouteSpec:
|
||||
route: SdkFunction
|
||||
rust_supported: bool = True
|
||||
|
||||
|
||||
TraceRouteSpec: TypeAlias = RouteSpec | GatewayRouteSpec
|
||||
|
|
@ -40,14 +66,7 @@ class TraceScenario:
|
|||
name: str
|
||||
fixture: Callable[[Engine, str], RouteFixture]
|
||||
mappings: tuple[TraceMapping, ...]
|
||||
modes: tuple[TraceMode, ...] = ("sync", "async")
|
||||
contract: TraceContract = TraceContract()
|
||||
sync_mappings: tuple[TraceMapping, ...] | None = None
|
||||
async_mappings: tuple[TraceMapping, ...] | None = None
|
||||
|
||||
def mappings_for(self, mode: TraceMode) -> tuple[TraceMapping, ...]:
|
||||
selected: Final = self.async_mappings if mode == "async" else self.sync_mappings
|
||||
return self.mappings if selected is None else selected
|
||||
asynchronous: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
|
|||
|
|
@ -1,31 +1,24 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from collections.abc import Sequence
|
||||
from typing import Final, Literal
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
from ...shared.reporting.models import SURFACES, CaseResult, RunStatus, SdkFunction, Surface
|
||||
from ...shared.reporting.rendering import ReportSection
|
||||
from ...shared.reporting.strategy import NotImplementedCaseSpec, SkippedCaseSpec
|
||||
from ...shared.tracing.steps import (
|
||||
PipelineStep,
|
||||
TraceContract,
|
||||
TraceDiff,
|
||||
TraceMapping,
|
||||
trace_depths,
|
||||
trace_diff,
|
||||
)
|
||||
from ...shared.tracing.steps import PipelineStep, trace_depths
|
||||
from .models import TraceEngine
|
||||
|
||||
TRACE_COMPARISON_ARTIFACT: Final = "trace_comparison"
|
||||
TRACE_ARTIFACT: Final = "trace"
|
||||
TRACE_PARITY_HINT: Final = (
|
||||
"rebuild the native bridge with the trace-parity feature, e.g. `uvx maturin develop --features trace-parity`"
|
||||
)
|
||||
|
||||
_COLORS: Final[dict[str, str]] = {"green": "32", "yellow": "33", "red": "31", "cyan": "36"}
|
||||
_COLORS: Final[dict[str, str]] = {"yellow": "33", "red": "31", "cyan": "36"}
|
||||
_RESET: Final = "\033[0m"
|
||||
|
||||
|
||||
|
|
@ -47,26 +40,15 @@ class TraceEventArtifact(BaseModel):
|
|||
return PipelineStep(self.id, self.parent_id, self.span, self.raw)
|
||||
|
||||
|
||||
class TraceMappingArtifact(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
span: str
|
||||
python: str | None
|
||||
rust: str | None
|
||||
|
||||
|
||||
class TraceComparisonArtifact(BaseModel):
|
||||
class TraceArtifact(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
engine: TraceEngine = "both"
|
||||
surface: Surface
|
||||
sdk_function: SdkFunction
|
||||
scenario: str
|
||||
mode: Literal["sync", "async"]
|
||||
mappings: tuple[TraceMappingArtifact, ...]
|
||||
python: tuple[TraceEventArtifact, ...]
|
||||
rust: tuple[TraceEventArtifact, ...]
|
||||
python_unmatched: int
|
||||
unordered_children_of: frozenset[str]
|
||||
python_error: str | None = None
|
||||
rust_error: str | None = None
|
||||
|
||||
|
|
@ -74,41 +56,27 @@ class TraceComparisonArtifact(BaseModel):
|
|||
def from_traces(
|
||||
cls,
|
||||
*,
|
||||
engine: TraceEngine = "both",
|
||||
surface: Surface,
|
||||
sdk_function: SdkFunction,
|
||||
scenario: str,
|
||||
mode: Literal["sync", "async"],
|
||||
mappings: Sequence[TraceMapping],
|
||||
contract: TraceContract,
|
||||
python: Sequence[PipelineStep],
|
||||
rust: Sequence[PipelineStep],
|
||||
python_unmatched: int,
|
||||
python_error: str | None = None,
|
||||
rust_error: str | None = None,
|
||||
) -> TraceComparisonArtifact:
|
||||
) -> TraceArtifact:
|
||||
return cls(
|
||||
engine=engine,
|
||||
surface=surface,
|
||||
sdk_function=sdk_function,
|
||||
scenario=scenario,
|
||||
mode=mode,
|
||||
mappings=tuple(
|
||||
TraceMappingArtifact(
|
||||
span=item.span,
|
||||
python=item.python.pattern if item.python else None,
|
||||
rust=item.rust,
|
||||
)
|
||||
for item in mappings
|
||||
),
|
||||
python=tuple(
|
||||
TraceEventArtifact(id=step.id, parent_id=step.parent_id, span=step.span, raw=step.raw)
|
||||
for step in python
|
||||
),
|
||||
rust=tuple(
|
||||
TraceEventArtifact(id=step.id, parent_id=step.parent_id, span=step.span, raw=step.raw)
|
||||
for step in rust
|
||||
TraceEventArtifact(id=step.id, parent_id=step.parent_id, span=step.span, raw=step.raw) for step in rust
|
||||
),
|
||||
python_unmatched=python_unmatched,
|
||||
unordered_children_of=contract.unordered_children_of,
|
||||
python_error=python_error,
|
||||
rust_error=rust_error,
|
||||
)
|
||||
|
|
@ -119,32 +87,9 @@ class TraceComparisonArtifact(BaseModel):
|
|||
def rust_steps(self) -> tuple[PipelineStep, ...]:
|
||||
return tuple(event.step() for event in self.rust)
|
||||
|
||||
def diff(self) -> TraceDiff:
|
||||
return trace_diff(
|
||||
self.python_steps(),
|
||||
self.rust_steps(),
|
||||
tuple(
|
||||
TraceMapping(
|
||||
item.span,
|
||||
re.compile(item.python) if item.python is not None else None,
|
||||
item.rust,
|
||||
)
|
||||
for item in self.mappings
|
||||
),
|
||||
TraceContract(self.unordered_children_of),
|
||||
)
|
||||
|
||||
def exact_match(self) -> bool:
|
||||
return self.diff().matches
|
||||
|
||||
def has_errors(self) -> bool:
|
||||
return self.python_error is not None or self.rust_error is not None
|
||||
|
||||
def contract_matches(self) -> bool:
|
||||
if self.has_errors():
|
||||
return False
|
||||
return self.diff().matches
|
||||
|
||||
|
||||
def _split_raw(raw: str) -> tuple[str, str]:
|
||||
location, separator, name = raw.partition(" ")
|
||||
|
|
@ -153,72 +98,28 @@ def _split_raw(raw: str) -> tuple[str, str]:
|
|||
return raw, ""
|
||||
|
||||
|
||||
def _python_line(index: int, step: PipelineStep, depth: int, exclusive: frozenset[str]) -> str:
|
||||
def _python_line(index: int, step: PipelineStep, depth: int) -> str:
|
||||
name: Final = _split_raw(step.raw)[0]
|
||||
location: Final = _split_raw(step.raw)[1]
|
||||
suffix: Final = f" ({location})" if location else ""
|
||||
marker: Final = " [python only]" if step.span in exclusive else ""
|
||||
return _paint(f"{index} {' ' * depth}{name}{suffix}{marker}", "cyan")
|
||||
return _paint(f"{index} {' ' * depth}{name}{suffix}", "cyan")
|
||||
|
||||
|
||||
def _python_lines(steps: tuple[PipelineStep, ...], exclusive: frozenset[str]) -> str:
|
||||
def _python_lines(steps: tuple[PipelineStep, ...]) -> str:
|
||||
depths: Final = trace_depths(steps)
|
||||
lines: Final = tuple(
|
||||
_python_line(index, step, depths[step.id], exclusive) for index, step in enumerate(steps, start=1)
|
||||
)
|
||||
lines: Final = tuple(_python_line(index, step, depths[step.id]) for index, step in enumerate(steps, start=1))
|
||||
return f"{_paint('PYTHON', 'cyan')} ({len(steps)} steps)\n" + ("\n".join(lines) if lines else "(empty)")
|
||||
|
||||
|
||||
def _python_references(steps: tuple[PipelineStep, ...]) -> dict[tuple[str, int], str]:
|
||||
references: dict[tuple[str, int], str] = {}
|
||||
occurrences: dict[str, int] = {}
|
||||
for index, step in enumerate(steps, start=1):
|
||||
name = _split_raw(step.raw)[0]
|
||||
occurrence = occurrences.get(step.span, 0) + 1
|
||||
occurrences[step.span] = occurrence
|
||||
references[(step.span, occurrence)] = f"{index} {name}"
|
||||
return references
|
||||
|
||||
|
||||
def _rust_line(
|
||||
step: PipelineStep,
|
||||
depth: int,
|
||||
occurrence: int,
|
||||
references: dict[tuple[str, int], str],
|
||||
) -> str:
|
||||
span: Final = _paint(step.span, "yellow")
|
||||
key: Final = (step.span, occurrence)
|
||||
reference: Final = (
|
||||
_paint(references[key], "cyan") if key in references else _paint("[rust only]", "yellow")
|
||||
)
|
||||
suffix: Final = f"#{occurrence}" if occurrence > 1 else ""
|
||||
return f"{' ' * depth}{span}{suffix} -> {reference}"
|
||||
|
||||
|
||||
def _rust_lines(steps: tuple[PipelineStep, ...], references: dict[tuple[str, int], str]) -> str:
|
||||
def _rust_lines(steps: tuple[PipelineStep, ...]) -> str:
|
||||
depths: Final = trace_depths(steps)
|
||||
occurrences: dict[str, int] = {}
|
||||
lines: list[str] = []
|
||||
for step in steps:
|
||||
occurrence = occurrences.get(step.span, 0) + 1
|
||||
occurrences[step.span] = occurrence
|
||||
lines.append(_rust_line(step, depths[step.id], occurrence, references))
|
||||
lines: Final = tuple(
|
||||
_paint(f"{index} {' ' * depths[step.id]}{step.span}", "yellow") for index, step in enumerate(steps, 1)
|
||||
)
|
||||
return f"{_paint('RUST', 'yellow')} ({len(steps)} steps)\n" + ("\n".join(lines) if lines else "(empty)")
|
||||
|
||||
|
||||
def _state_text(state: str, *, good: bool) -> str:
|
||||
return _paint(state, "green" if good else "red")
|
||||
|
||||
|
||||
def _contract_line(artifact: TraceComparisonArtifact) -> str:
|
||||
matches: Final = artifact.contract_matches()
|
||||
status: Final = _state_text("PASS" if matches else "FAIL", good=matches)
|
||||
if artifact.python_error or artifact.rust_error:
|
||||
return f"Contract: {status}"
|
||||
return f"Contract: {status}"
|
||||
|
||||
|
||||
def _error_lines(artifact: TraceComparisonArtifact) -> tuple[str, ...]:
|
||||
def _error_lines(artifact: TraceArtifact) -> tuple[str, ...]:
|
||||
lines: list[str] = []
|
||||
for engine, error in (("Python", artifact.python_error), ("Rust", artifact.rust_error)):
|
||||
if error is None:
|
||||
|
|
@ -229,68 +130,20 @@ def _error_lines(artifact: TraceComparisonArtifact) -> tuple[str, ...]:
|
|||
return tuple(lines)
|
||||
|
||||
|
||||
def _unseen_mappings(
|
||||
artifact: TraceComparisonArtifact,
|
||||
python: tuple[PipelineStep, ...],
|
||||
rust: tuple[PipelineStep, ...],
|
||||
) -> tuple[str, ...]:
|
||||
return artifact.diff().missing_mappings
|
||||
|
||||
|
||||
def _comparison_status_lines(
|
||||
artifact: TraceComparisonArtifact,
|
||||
python: tuple[PipelineStep, ...],
|
||||
rust: tuple[PipelineStep, ...],
|
||||
) -> tuple[str, ...]:
|
||||
diff: Final = artifact.diff()
|
||||
exact_match: Final = artifact.exact_match()
|
||||
if artifact.has_errors():
|
||||
return (*_error_lines(artifact), _contract_line(artifact))
|
||||
unseen: Final = _unseen_mappings(artifact, python, rust)
|
||||
unseen_line: Final[tuple[str, ...]] = (f"Unseen mappings: {', '.join(unseen)}",) if unseen else ()
|
||||
drift_lines: Final[tuple[str, ...]] = (
|
||||
(_state_text("Same steps, order, and nesting", good=True),)
|
||||
if exact_match
|
||||
else (
|
||||
_paint(f"Python only: {', '.join(diff.python_only) or 'none'}", "cyan"),
|
||||
_paint(f"Rust only: {', '.join(diff.rust_only) or 'none'}", "yellow"),
|
||||
f"First difference: {diff.first_difference or 'none'}",
|
||||
f"Python frames outside mapping: {artifact.python_unmatched}",
|
||||
)
|
||||
)
|
||||
return (
|
||||
f"Trace: {_state_text('MATCH' if exact_match else 'DRIFT', good=exact_match)}",
|
||||
*drift_lines,
|
||||
*unseen_line,
|
||||
_contract_line(artifact),
|
||||
)
|
||||
|
||||
|
||||
def _render_comparison(artifact: TraceComparisonArtifact) -> str:
|
||||
python: Final = artifact.python_steps()
|
||||
rust: Final = artifact.rust_steps()
|
||||
diff: Final = artifact.diff()
|
||||
python_exclusive: Final = frozenset(item.span for item in artifact.mappings if item.rust is None)
|
||||
status_lines: Final = _comparison_status_lines(artifact, python, rust)
|
||||
return "\n\n".join(
|
||||
(
|
||||
_python_lines(python, python_exclusive | frozenset(diff.python_only)),
|
||||
_rust_lines(rust, _python_references(python)),
|
||||
"\n".join(status_lines),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _mode(nodeid: str) -> str:
|
||||
if "[" in nodeid:
|
||||
return nodeid.rsplit("[", 1)[-1].removesuffix("]")
|
||||
head, _, tail = nodeid.rpartition(":")
|
||||
return tail if head else "unknown mode"
|
||||
def _render_trace(artifact: TraceArtifact) -> str:
|
||||
traces: tuple[str, ...]
|
||||
if artifact.engine == "python":
|
||||
traces = (_python_lines(artifact.python_steps()),)
|
||||
elif artifact.engine == "rust":
|
||||
traces = (_rust_lines(artifact.rust_steps()),)
|
||||
else:
|
||||
traces = (_python_lines(artifact.python_steps()), _rust_lines(artifact.rust_steps()))
|
||||
return "\n\n".join((*traces, *_error_lines(artifact)))
|
||||
|
||||
|
||||
def _scenario(nodeid: str) -> str:
|
||||
parts: Final = nodeid.split(":")
|
||||
return parts[-2] if len(parts) >= 5 else "default"
|
||||
return parts[-1] if len(parts) >= 4 else "default"
|
||||
|
||||
|
||||
def _unavailable(status: RunStatus) -> str:
|
||||
|
|
@ -299,20 +152,20 @@ def _unavailable(status: RunStatus) -> str:
|
|||
|
||||
def _render_artifact(body: str) -> str:
|
||||
try:
|
||||
artifact: Final = TraceComparisonArtifact.model_validate_json(body)
|
||||
artifact: Final = TraceArtifact.model_validate_json(body)
|
||||
except ValidationError as error:
|
||||
return f"Trace comparison artifact is invalid: {error}"
|
||||
return _render_comparison(artifact)
|
||||
return f"Trace artifact is invalid: {error}"
|
||||
return _render_trace(artifact)
|
||||
|
||||
|
||||
def _mode_section(result: CaseResult, nodeid: str, status: RunStatus) -> str:
|
||||
def _scenario_section(result: CaseResult, nodeid: str, status: RunStatus) -> str:
|
||||
artifacts: Final = tuple(
|
||||
artifact for artifact in result.artifacts.get(nodeid, ()) if artifact.kind == TRACE_COMPARISON_ARTIFACT
|
||||
artifact for artifact in result.artifacts.get(nodeid, ()) if artifact.kind == TRACE_ARTIFACT
|
||||
)
|
||||
body: Final = (
|
||||
"\n\n".join(_render_artifact(artifact.body) for artifact in artifacts) if artifacts else _unavailable(status)
|
||||
)
|
||||
label: Final = f"Scenario: {_scenario(nodeid)} / Mode: {_mode(nodeid)}"
|
||||
label: Final = f"Scenario: {_scenario(nodeid)}"
|
||||
return f"{label}\n{'-' * len(label)}\n\n{body}"
|
||||
|
||||
|
||||
|
|
@ -321,7 +174,7 @@ def _case_block(result: CaseResult) -> str:
|
|||
outcomes: Final = tuple(result.outcomes.items()) or (
|
||||
(nodeid, RunStatus.NOT_RUN) for nodeid in sorted(result.collected)
|
||||
)
|
||||
sections: Final = tuple(_mode_section(result, nodeid, status) for nodeid, status in outcomes)
|
||||
sections: Final = tuple(_scenario_section(result, nodeid, status) for nodeid, status in outcomes)
|
||||
return "\n\n".join((f"{header}\n{'=' * len(header)}", *sections))
|
||||
|
||||
|
||||
|
|
@ -357,11 +210,11 @@ def _surface_section(surface: Surface, results: Sequence[CaseResult]) -> ReportS
|
|||
*((not_implemented,) if not_implemented else ()),
|
||||
*((skipped,) if skipped else ()),
|
||||
)
|
||||
return ReportSection(f"{surface.upper()} trace comparisons", blocks or ("No runnable trace comparisons",))
|
||||
return ReportSection(f"{surface.upper()} traces", blocks or ("No runnable traces",))
|
||||
|
||||
|
||||
def render_trace_results(results: Sequence[CaseResult]) -> tuple[ReportSection, ...]:
|
||||
sections: Final = tuple(
|
||||
section for surface in SURFACES if (section := _surface_section(surface, results)) is not None
|
||||
)
|
||||
return sections or (ReportSection("Trace comparisons", ("No trace comparisons selected",)),)
|
||||
return sections or (ReportSection("Traces", ("No traces selected",)),)
|
||||
|
|
|
|||
|
|
@ -4,13 +4,20 @@ import importlib
|
|||
from collections.abc import Sequence
|
||||
from pathlib import Path
|
||||
from time import monotonic
|
||||
from typing import Final
|
||||
from typing import Final, cast
|
||||
|
||||
from ...shared.native_build import ensure_trace_bridge
|
||||
from ...shared.reporting.models import CaseResult, HarnessCase, HarnessRun, ResultArtifact, RunStatus, Surface
|
||||
from ...shared.reporting.strategy import ModuleCaseSpec, UpdateCallback
|
||||
from ...shared.native_build import ensure_trace_bridge
|
||||
from .models import GatewayRouteSpec, RouteSpec, TraceExecutionFailure, TraceMode, TraceScenario, TraceSuite
|
||||
from .reporting import TRACE_COMPARISON_ARTIFACT, TraceComparisonArtifact
|
||||
from .models import (
|
||||
GatewayRouteSpec,
|
||||
RouteSpec,
|
||||
TraceEngine,
|
||||
TraceExecutionFailure,
|
||||
TraceScenario,
|
||||
TraceSuite,
|
||||
)
|
||||
from .reporting import TRACE_ARTIFACT, TraceArtifact
|
||||
from .sdk.execution import execute_trace
|
||||
|
||||
|
||||
|
|
@ -32,15 +39,13 @@ def validate_trace_suite(suite: TraceSuite, harness_case: HarnessCase) -> str |
|
|||
names: Final = tuple(scenario.name for scenario in suite.scenarios)
|
||||
if not names or len(names) != len(set(names)) or any(not name or ":" in name for name in names):
|
||||
return "scenario names must be non-empty, unique, and colon-free"
|
||||
invalid_modes: Final = tuple(
|
||||
invalid_names: Final = tuple(
|
||||
scenario.name
|
||||
for scenario in suite.scenarios
|
||||
if not scenario.modes
|
||||
or len(scenario.modes) != len(set(scenario.modes))
|
||||
or any(mode not in {"sync", "async"} for mode in scenario.modes)
|
||||
if not scenario.name.startswith("async-" if scenario.asynchronous else "sync-")
|
||||
)
|
||||
if invalid_modes:
|
||||
return f"scenarios must use non-empty, unique sync/async modes: {', '.join(invalid_modes)}"
|
||||
if invalid_names:
|
||||
return f"scenario names must start with sync- or async-: {', '.join(invalid_names)}"
|
||||
surface: Final = harness_case.surface
|
||||
if surface == "sdk" and not isinstance(suite.route, RouteSpec):
|
||||
return "must use RouteSpec for the sdk surface"
|
||||
|
|
@ -57,15 +62,14 @@ def scenario_nodeids(
|
|||
trace_suite: TraceSuite,
|
||||
harness_case: HarnessCase,
|
||||
selected_scenarios: frozenset[str] = frozenset(),
|
||||
) -> tuple[tuple[TraceScenario, TraceMode, str], ...]:
|
||||
) -> tuple[tuple[TraceScenario, str], ...]:
|
||||
surface: Final = harness_case.surface
|
||||
if surface is None:
|
||||
return ()
|
||||
return tuple(
|
||||
(scenario, mode, f"trace:{surface}:{harness_case.sdk_function}:{scenario.name}:{mode}")
|
||||
(scenario, f"trace:{surface}:{harness_case.sdk_function}:{scenario.name}")
|
||||
for scenario in trace_suite.scenarios
|
||||
if not selected_scenarios or scenario.name in selected_scenarios
|
||||
for mode in scenario.modes
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -77,54 +81,49 @@ def _record_setup_failure(run: HarnessRun, case: HarnessCase, message: str, stag
|
|||
run.failures.append((nodeid, message))
|
||||
|
||||
|
||||
def run_trace_mode(
|
||||
def run_trace_scenario(
|
||||
run: HarnessRun,
|
||||
result: CaseResult,
|
||||
trace_suite: TraceSuite,
|
||||
scenario: TraceScenario,
|
||||
mode: TraceMode,
|
||||
surface: Surface,
|
||||
nodeid: str,
|
||||
on_update: UpdateCallback,
|
||||
engine: TraceEngine = "both",
|
||||
) -> None:
|
||||
started_at: Final = monotonic()
|
||||
comparison: Final = _execute_mode(trace_suite, scenario, mode, surface)
|
||||
trace: Final = _execute_scenario(trace_suite, scenario, surface, engine)
|
||||
duration: Final = monotonic() - started_at
|
||||
if isinstance(comparison, TraceExecutionFailure):
|
||||
if isinstance(trace, TraceExecutionFailure):
|
||||
result.record(nodeid, RunStatus.ERROR, duration)
|
||||
run.failures.append((nodeid, comparison.message))
|
||||
run.failures.append((nodeid, trace.message))
|
||||
on_update(run)
|
||||
return
|
||||
artifact: Final = ResultArtifact(TRACE_COMPARISON_ARTIFACT, comparison.model_dump_json())
|
||||
if comparison.has_errors():
|
||||
artifact: Final = ResultArtifact(TRACE_ARTIFACT, trace.model_dump_json())
|
||||
if trace.has_errors():
|
||||
result.record(nodeid, RunStatus.ERROR, duration, (artifact,))
|
||||
run.failures.append(
|
||||
(nodeid, "\n".join(error for error in (comparison.python_error, comparison.rust_error) if error))
|
||||
)
|
||||
run.failures.append((nodeid, "\n".join(error for error in (trace.python_error, trace.rust_error) if error)))
|
||||
else:
|
||||
status: Final = RunStatus.PASSED if comparison.contract_matches() else RunStatus.FAILED
|
||||
result.record(nodeid, status, duration, (artifact,))
|
||||
if status is RunStatus.FAILED:
|
||||
run.failures.append((nodeid, "trace contract mismatch; see the rendered comparison"))
|
||||
result.record(nodeid, RunStatus.PASSED, duration, (artifact,))
|
||||
on_update(run)
|
||||
|
||||
|
||||
def _execute_mode(
|
||||
def _execute_scenario(
|
||||
trace_suite: TraceSuite,
|
||||
scenario: TraceScenario,
|
||||
mode: TraceMode,
|
||||
surface: Surface,
|
||||
) -> TraceComparisonArtifact | TraceExecutionFailure:
|
||||
engine: TraceEngine,
|
||||
) -> TraceArtifact | TraceExecutionFailure:
|
||||
route: Final = trace_suite.route
|
||||
if isinstance(route, GatewayRouteSpec):
|
||||
if surface != "gateway":
|
||||
return TraceExecutionFailure("harness", "gateway route cannot run on the sdk surface")
|
||||
from .gateway.execution import execute_gateway_trace
|
||||
|
||||
return execute_gateway_trace(route, scenario, mode)
|
||||
return execute_gateway_trace(route, scenario, engine)
|
||||
if surface != "sdk":
|
||||
return TraceExecutionFailure("harness", "sdk route cannot run on the gateway surface")
|
||||
return execute_trace(route, scenario, mode, surface)
|
||||
return execute_trace(route, scenario, surface, engine)
|
||||
|
||||
|
||||
def _run_case(
|
||||
|
|
@ -132,6 +131,7 @@ def _run_case(
|
|||
harness_case: HarnessCase,
|
||||
selected_scenarios: frozenset[str],
|
||||
on_update: UpdateCallback,
|
||||
engine: TraceEngine,
|
||||
) -> None:
|
||||
result: Final = run.results[harness_case.key]
|
||||
spec: Final = harness_case.spec
|
||||
|
|
@ -146,15 +146,29 @@ def _run_case(
|
|||
on_update(run)
|
||||
return
|
||||
nodeids: Final = scenario_nodeids(trace_suite, harness_case, selected_scenarios)
|
||||
result.collected.update(nodeid for _, _, nodeid in nodeids)
|
||||
result.collected.update(nodeid for _, nodeid in nodeids)
|
||||
if not nodeids:
|
||||
result.status = RunStatus.SKIPPED
|
||||
on_update(run)
|
||||
return
|
||||
result.status = RunStatus.RUNNING
|
||||
on_update(run)
|
||||
for scenario, mode, nodeid in nodeids:
|
||||
run_trace_mode(run, result, trace_suite, scenario, mode, surface, nodeid, on_update)
|
||||
for scenario, nodeid in nodeids:
|
||||
run_trace_scenario(run, result, trace_suite, scenario, surface, nodeid, on_update, engine)
|
||||
|
||||
|
||||
def runner_selection(runner_args: Sequence[str]) -> tuple[frozenset[str], TraceEngine]:
|
||||
engine: TraceEngine = "both"
|
||||
scenarios: list[str] = []
|
||||
for argument in runner_args:
|
||||
if argument.startswith("--engine="):
|
||||
value = argument.removeprefix("--engine=")
|
||||
if value not in {"python", "rust"}:
|
||||
raise ValueError(f"invalid trace engine: {value}")
|
||||
engine = cast(TraceEngine, value)
|
||||
else:
|
||||
scenarios.append(argument)
|
||||
return frozenset(scenarios), engine
|
||||
|
||||
|
||||
def run_trace_cases(
|
||||
|
|
@ -163,10 +177,10 @@ def run_trace_cases(
|
|||
on_update: UpdateCallback,
|
||||
runner_args: Sequence[str] = (),
|
||||
) -> tuple[int, HarnessRun]:
|
||||
selected_scenarios: Final = frozenset(runner_args)
|
||||
selected_scenarios, engine = runner_selection(runner_args)
|
||||
run: Final = HarnessRun.from_cases(cases)
|
||||
runnable_cases: Final = tuple(case for case in cases if isinstance(case.spec, ModuleCaseSpec))
|
||||
bridge_error: Final = ensure_trace_bridge(repo_root) if runnable_cases else None
|
||||
bridge_error: Final = ensure_trace_bridge(repo_root) if runnable_cases and engine != "python" else None
|
||||
if bridge_error is not None:
|
||||
for harness_case in runnable_cases:
|
||||
_record_setup_failure(run, harness_case, bridge_error, "bridge")
|
||||
|
|
@ -174,7 +188,7 @@ def run_trace_cases(
|
|||
on_update(run)
|
||||
return 1, run
|
||||
for harness_case in cases:
|
||||
_run_case(run, harness_case, selected_scenarios, on_update)
|
||||
_run_case(run, harness_case, selected_scenarios, on_update, engine)
|
||||
run.finished_at = monotonic()
|
||||
on_update(run)
|
||||
failed: Final = any(
|
||||
|
|
|
|||
|
|
@ -1,10 +1,15 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
from .....shared.parity.recorded_http import HttpHeader, RecordedHttpResponse
|
||||
from .....shared.tracing.steps import Engine, mapping
|
||||
from ...fixtures import (
|
||||
anthropic_response_body,
|
||||
anthropic_stream_events,
|
||||
aws_event_stream_response,
|
||||
json_response,
|
||||
sse_response,
|
||||
)
|
||||
from ...models import RouteFixture, RouteSpec, TraceScenario, TraceSuite
|
||||
|
||||
COMMON_MAPPINGS: Final = (
|
||||
|
|
@ -24,6 +29,22 @@ COMMON_MAPPINGS: Final = (
|
|||
mapping(rust_span="execute_chat_completions_provider_call"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(rust_span="transform_response", python_frame=r"(?<!async_)transform_response$"),
|
||||
mapping(span="python_logging_pre_call", python_frame=r"Logging\.pre_call$"),
|
||||
mapping(span="python_logging_post_call", python_frame=r"Logging\.post_call$"),
|
||||
mapping(span="python_success_callback", python_frame=r"Logging\.async_success_handler$|Logging\.success_handler$"),
|
||||
)
|
||||
|
||||
STREAM_MAPPINGS: Final = (
|
||||
mapping(span="python_stream_wrapper", python_frame=r"CustomStreamWrapper\.__init__$"),
|
||||
mapping(span="python_stream_next", python_frame=r"CustomStreamWrapper\.__next__$|CustomStreamWrapper\.__anext__$"),
|
||||
mapping(span="python_stream_chunk", python_frame=r"CustomStreamWrapper\.chunk_creator$"),
|
||||
mapping(span="python_stream_finalize", python_frame=r"CustomStreamWrapper\._finalize_completed_stream$"),
|
||||
)
|
||||
|
||||
FAILURE_MAPPINGS: Final = (
|
||||
mapping(span="python_exception_mapping", python_frame=r"(?<!_)exception_type$"),
|
||||
mapping(span="python_failure_callback", python_frame=r"Logging\.failure_handler$"),
|
||||
mapping(span="python_async_failure_callback", python_frame=r"Logging\.async_failure_handler$"),
|
||||
)
|
||||
|
||||
SYNC_MAPPINGS: Final = (
|
||||
|
|
@ -44,41 +65,23 @@ ASYNC_MAPPINGS: Final = (
|
|||
|
||||
|
||||
def _anthropic_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
response: Final = json.dumps(
|
||||
{
|
||||
"id": "msg_trace",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-5",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 3},
|
||||
}
|
||||
).encode()
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
"model": "anthropic/claude-sonnet-5",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
**({"optional_params": {"max_tokens": 16}} if engine == "rust" else {"max_tokens": 16}),
|
||||
},
|
||||
provider_responses=(
|
||||
RecordedHttpResponse.from_bytes(
|
||||
200, (HttpHeader(name="content-type", value="application/json"),), response
|
||||
),
|
||||
),
|
||||
provider_responses=(json_response(anthropic_response_body()),),
|
||||
)
|
||||
|
||||
|
||||
def _bedrock_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
response: Final = json.dumps(
|
||||
{
|
||||
"output": {"message": {"role": "assistant", "content": [{"text": "hello"}]}},
|
||||
"stopReason": "end_turn",
|
||||
"usage": {"inputTokens": 2, "outputTokens": 3, "totalTokens": 5},
|
||||
"metrics": {"latencyMs": 1},
|
||||
}
|
||||
).encode()
|
||||
response: Final[dict[str, object]] = {
|
||||
"output": {"message": {"role": "assistant", "content": [{"text": "hello"}]}},
|
||||
"stopReason": "end_turn",
|
||||
"usage": {"inputTokens": 2, "outputTokens": 3, "totalTokens": 5},
|
||||
"metrics": {"latencyMs": 1},
|
||||
}
|
||||
credentials: Final = {
|
||||
"aws_access_key_id": "test-access",
|
||||
"aws_secret_access_key": "test-secret",
|
||||
|
|
@ -94,11 +97,60 @@ def _bedrock_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
|||
else {**credentials, "max_tokens": 16}
|
||||
),
|
||||
},
|
||||
provider_responses=(json_response(response),),
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_stream_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _anthropic_fixture(engine, _base_url)
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(sse_response(anthropic_stream_events()),),
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
def _bedrock_stream_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _bedrock_fixture(engine, _base_url)
|
||||
events: Final[tuple[dict[str, object], ...]] = (
|
||||
{"messageStart": {"role": "assistant"}},
|
||||
{"contentBlockStart": {"contentBlockIndex": 0, "start": {}}},
|
||||
{"contentBlockDelta": {"contentBlockIndex": 0, "delta": {"text": "hello"}}},
|
||||
{"contentBlockStop": {"contentBlockIndex": 0}},
|
||||
{"messageStop": {"stopReason": "end_turn"}},
|
||||
{"metadata": {"usage": {"inputTokens": 2, "outputTokens": 1, "totalTokens": 3}}},
|
||||
)
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(aws_event_stream_response(events),),
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
def _provider_error_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _anthropic_fixture(engine, _base_url)
|
||||
return fixture.derive(
|
||||
provider_responses=(
|
||||
RecordedHttpResponse.from_bytes(
|
||||
200, (HttpHeader(name="content-type", value="application/json"),), response
|
||||
json_response(
|
||||
{"type": "error", "error": {"type": "invalid_request_error", "message": "bad request"}},
|
||||
status=400,
|
||||
),
|
||||
),
|
||||
expected_failure=True,
|
||||
)
|
||||
|
||||
|
||||
def _stream_error_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
fixture: Final = _anthropic_fixture(engine, base_url)
|
||||
events: Final = (
|
||||
anthropic_stream_events()[0],
|
||||
("error", {"type": "error", "error": {"type": "overloaded_error", "message": "overloaded"}}),
|
||||
)
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(sse_response(events),),
|
||||
expected_failure=True,
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -132,18 +184,64 @@ TRACE_SUITE: Final = TraceSuite(
|
|||
route=SPEC,
|
||||
scenarios=(
|
||||
TraceScenario(
|
||||
name="anthropic",
|
||||
name="sync-anthropic",
|
||||
fixture=_anthropic_fixture,
|
||||
mappings=COMMON_MAPPINGS,
|
||||
sync_mappings=SYNC_MAPPINGS,
|
||||
async_mappings=ASYNC_MAPPINGS,
|
||||
mappings=SYNC_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="bedrock",
|
||||
name="async-anthropic",
|
||||
fixture=_anthropic_fixture,
|
||||
mappings=ASYNC_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-anthropic-stream",
|
||||
fixture=_anthropic_stream_fixture,
|
||||
mappings=(*SYNC_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-anthropic-stream",
|
||||
fixture=_anthropic_stream_fixture,
|
||||
mappings=(*ASYNC_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-anthropic-provider-error",
|
||||
fixture=_provider_error_fixture,
|
||||
mappings=(*ASYNC_MAPPINGS, *FAILURE_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-anthropic-stream-error",
|
||||
fixture=_stream_error_fixture,
|
||||
mappings=(*ASYNC_MAPPINGS, *STREAM_MAPPINGS, *FAILURE_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-bedrock",
|
||||
fixture=_bedrock_fixture,
|
||||
mappings=BEDROCK_COMMON_MAPPINGS,
|
||||
sync_mappings=BEDROCK_SYNC_MAPPINGS,
|
||||
async_mappings=BEDROCK_ASYNC_MAPPINGS,
|
||||
mappings=BEDROCK_SYNC_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-bedrock",
|
||||
fixture=_bedrock_fixture,
|
||||
mappings=BEDROCK_ASYNC_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-bedrock-event-stream",
|
||||
fixture=_bedrock_stream_fixture,
|
||||
mappings=(*BEDROCK_SYNC_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-bedrock-event-stream",
|
||||
fixture=_bedrock_stream_fixture,
|
||||
mappings=(*BEDROCK_ASYNC_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,18 +1,20 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable
|
||||
import os
|
||||
from collections.abc import AsyncIterable, Awaitable, Iterable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final, Protocol, cast
|
||||
from unittest.mock import patch
|
||||
|
||||
from ....shared.parity.replay import replay_server
|
||||
from ....shared.reporting.models import Surface
|
||||
from ....shared.tracing.native import TraceResponsePayload, native_trace_events
|
||||
from ....shared.tracing.profiler import FunctionTraceEvent, profile_python
|
||||
from ....shared.tracing.steps import Engine, pipeline_projection
|
||||
from ..models import RouteFixture, RouteSpec, TraceExecutionFailure, TraceMode, TraceScenario
|
||||
from ..reporting import TraceComparisonArtifact
|
||||
from ..models import RouteFixture, RouteSpec, TraceEngine, TraceExecutionFailure, TraceScenario
|
||||
from ..reporting import TraceArtifact
|
||||
|
||||
|
||||
class SdkCall(Protocol):
|
||||
|
|
@ -25,10 +27,20 @@ class _CollectedTrace:
|
|||
error: str | None = None
|
||||
|
||||
|
||||
def _invoke(function: SdkCall, kwargs: dict[str, object], *, asynchronous: bool) -> object:
|
||||
def _invoke(
|
||||
function: SdkCall,
|
||||
kwargs: dict[str, object],
|
||||
*,
|
||||
asynchronous: bool,
|
||||
consume_stream: bool = False,
|
||||
) -> object:
|
||||
async def invoke_async() -> object:
|
||||
try:
|
||||
return await cast(Awaitable[object], function(**kwargs))
|
||||
response: Final = await cast(Awaitable[object], function(**kwargs))
|
||||
if consume_stream and isinstance(response, AsyncIterable):
|
||||
stream = cast(AsyncIterable[object], response)
|
||||
return tuple([item async for item in stream])
|
||||
return response
|
||||
finally:
|
||||
await asyncio.sleep(0)
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
|
@ -38,7 +50,10 @@ def _invoke(function: SdkCall, kwargs: dict[str, object], *, asynchronous: bool)
|
|||
|
||||
if asynchronous:
|
||||
return asyncio.run(invoke_async())
|
||||
return function(**kwargs)
|
||||
response: Final = function(**kwargs)
|
||||
if consume_stream and isinstance(response, Iterable):
|
||||
return tuple(cast(Iterable[object], response))
|
||||
return response
|
||||
|
||||
|
||||
def _entrypoint(spec: RouteSpec, engine: Engine, *, asynchronous: bool) -> SdkCall | TraceExecutionFailure:
|
||||
|
|
@ -47,6 +62,8 @@ def _entrypoint(spec: RouteSpec, engine: Engine, *, asynchronous: bool) -> SdkCa
|
|||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
if engine == "rust":
|
||||
if spec.rust_entrypoints is None:
|
||||
return TraceExecutionFailure("rust", f"{spec.route} has no native Rust trace entrypoint")
|
||||
bridge: Final = cast(object | None, get_native_bridge())
|
||||
if bridge is None:
|
||||
return TraceExecutionFailure("rust", "native Rust bridge is required for trace parity")
|
||||
|
|
@ -62,14 +79,6 @@ def _entrypoint(spec: RouteSpec, engine: Engine, *, asynchronous: bool) -> SdkCa
|
|||
return cast(SdkCall, getattr(owner, spec.python_entrypoints[int(asynchronous)]))
|
||||
|
||||
|
||||
def _python_invocation_error(function: SdkCall, kwargs: dict[str, object], *, asynchronous: bool) -> str | None:
|
||||
try:
|
||||
_invoke(function, kwargs, asynchronous=asynchronous)
|
||||
except Exception as error:
|
||||
return f"{type(error).__name__}: {error}"
|
||||
return None
|
||||
|
||||
|
||||
def _collect(
|
||||
function: SdkCall,
|
||||
fixture: RouteFixture,
|
||||
|
|
@ -83,14 +92,23 @@ def _collect(
|
|||
return _CollectedTrace(native_trace_events(payload), payload.error)
|
||||
import litellm
|
||||
|
||||
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
|
||||
error: Final = _python_invocation_error(function, kwargs, asynchronous=asynchronous)
|
||||
previous_suppress_debug_info: Final = litellm.suppress_debug_info
|
||||
try:
|
||||
if fixture.expected_failure:
|
||||
litellm.suppress_debug_info = True
|
||||
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
|
||||
error: str | None
|
||||
try:
|
||||
_invoke(function, kwargs, asynchronous=asynchronous, consume_stream=fixture.consume_stream)
|
||||
error = None
|
||||
except Exception as caught:
|
||||
error = f"{type(caught).__name__}: {caught}"
|
||||
finally:
|
||||
litellm.suppress_debug_info = previous_suppress_debug_info
|
||||
return _CollectedTrace(tuple(profiler.events), error)
|
||||
|
||||
|
||||
def collect_trace(
|
||||
spec: RouteSpec, engine: Engine, *, asynchronous: bool
|
||||
) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
|
||||
def collect_trace(spec: RouteSpec, engine: Engine, *, asynchronous: bool) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
|
||||
function: Final = _entrypoint(spec, engine, asynchronous=asynchronous)
|
||||
if isinstance(function, TraceExecutionFailure):
|
||||
return function
|
||||
|
|
@ -101,15 +119,18 @@ def collect_trace(
|
|||
provider.enqueue_response(response)
|
||||
fixture: Final = RouteFixture(
|
||||
kwargs={
|
||||
**base_fixture.kwargs,
|
||||
"api_key": "test-key",
|
||||
**base_fixture.kwargs,
|
||||
"api_base": provider.url,
|
||||
**({"timeout_seconds": 5} if engine == "rust" else {"timeout": 5}),
|
||||
},
|
||||
provider_responses=base_fixture.provider_responses,
|
||||
expected_failure=base_fixture.expected_failure,
|
||||
consume_stream=base_fixture.consume_stream,
|
||||
environment=base_fixture.environment,
|
||||
)
|
||||
collected: Final = _collect(function, fixture, engine, asynchronous=asynchronous)
|
||||
with patch.dict(os.environ, fixture.environment):
|
||||
collected: Final = _collect(function, fixture, engine, asynchronous=asynchronous)
|
||||
provider.take_requests(len(fixture.provider_responses))
|
||||
except Exception as error:
|
||||
return TraceExecutionFailure(engine, f"{type(error).__name__}: {error}")
|
||||
|
|
@ -129,48 +150,56 @@ def _failure_message(result: tuple[FunctionTraceEvent, ...] | TraceExecutionFail
|
|||
|
||||
|
||||
def execute_trace(
|
||||
route: RouteSpec, scenario: TraceScenario, mode: TraceMode, surface: Surface
|
||||
) -> TraceComparisonArtifact:
|
||||
asynchronous: Final = mode == "async"
|
||||
mappings: Final = scenario.mappings_for(mode)
|
||||
route: RouteSpec,
|
||||
scenario: TraceScenario,
|
||||
surface: Surface,
|
||||
engine: TraceEngine = "both",
|
||||
) -> TraceArtifact:
|
||||
effective_engine: Final[TraceEngine] = "python" if engine == "both" and route.rust_entrypoints is None else engine
|
||||
scenario_route: Final = RouteSpec(
|
||||
route=route.route,
|
||||
python_entrypoints=route.python_entrypoints,
|
||||
rust_entrypoints=route.rust_entrypoints,
|
||||
fixture=scenario.fixture,
|
||||
)
|
||||
python_trace: Final = collect_trace(scenario_route, "python", asynchronous=asynchronous)
|
||||
rust_trace: Final = collect_trace(scenario_route, "rust", asynchronous=asynchronous)
|
||||
python_trace: Final = (
|
||||
collect_trace(
|
||||
scenario_route,
|
||||
"python",
|
||||
asynchronous=scenario.asynchronous,
|
||||
)
|
||||
if effective_engine != "rust"
|
||||
else ()
|
||||
)
|
||||
rust_trace: Final = (
|
||||
collect_trace(scenario_route, "rust", asynchronous=scenario.asynchronous)
|
||||
if effective_engine != "python"
|
||||
else ()
|
||||
)
|
||||
python_error: Final = _failure_message(python_trace)
|
||||
rust_error: Final = _failure_message(rust_trace)
|
||||
python_events: Final = python_trace if isinstance(python_trace, tuple) else ()
|
||||
rust_events: Final = rust_trace if isinstance(rust_trace, tuple) else ()
|
||||
try:
|
||||
python: Final = pipeline_projection("python", python_events, mappings)
|
||||
rust: Final = pipeline_projection("rust", rust_events, mappings)
|
||||
python: Final = pipeline_projection("python", python_events)
|
||||
rust: Final = pipeline_projection("rust", rust_events)
|
||||
except ValueError as error:
|
||||
return TraceComparisonArtifact.from_traces(
|
||||
return TraceArtifact.from_traces(
|
||||
engine=effective_engine,
|
||||
surface=surface,
|
||||
sdk_function=route.route,
|
||||
scenario=scenario.name,
|
||||
mode=mode,
|
||||
mappings=mappings,
|
||||
contract=scenario.contract,
|
||||
python=(),
|
||||
rust=(),
|
||||
python_unmatched=0,
|
||||
python_error=f"harness: {error}",
|
||||
)
|
||||
return TraceComparisonArtifact.from_traces(
|
||||
return TraceArtifact.from_traces(
|
||||
engine=effective_engine,
|
||||
surface=surface,
|
||||
sdk_function=route.route,
|
||||
scenario=scenario.name,
|
||||
mode=mode,
|
||||
mappings=mappings,
|
||||
contract=scenario.contract,
|
||||
python=python.steps,
|
||||
rust=rust.steps,
|
||||
python_unmatched=python.unmatched,
|
||||
python_error=python_error,
|
||||
rust_error=rust_error,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,14 +1,26 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
from .....shared.parity.recorded_http import HttpHeader, RecordedHttpResponse
|
||||
from .....shared.tracing.steps import Engine, mapping
|
||||
from ...fixtures import (
|
||||
anthropic_response_body,
|
||||
anthropic_stream_events,
|
||||
aws_event_stream_response,
|
||||
json_response,
|
||||
sse_response,
|
||||
)
|
||||
from ...models import RouteFixture, RouteSpec, TraceScenario, TraceSuite
|
||||
|
||||
COMMON_MAPPINGS: Final = (
|
||||
mapping(rust_span="messages", python_frame=r"anthropic_interface/messages/__init__\.py:\d+ a?create$"),
|
||||
mapping(span="python_sanitize_empty_content", python_frame=r"strip_empty_content_blocks_from_anthropic_messages$"),
|
||||
mapping(span="python_sanitize_tool_ids", python_frame=r"sanitize_tool_use_ids_in_anthropic_messages$"),
|
||||
mapping(
|
||||
span="python_flatten_web_search", python_frame=r"flatten_unencrypted_web_search_results_in_anthropic_messages$"
|
||||
),
|
||||
mapping(span="python_cache_control", python_frame=r"AnthropicCacheControlHook\.maybe_inject_cache_control$"),
|
||||
mapping(span="python_pre_request_hooks", python_frame=r"_execute_pre_request_hooks$"),
|
||||
mapping(
|
||||
span="python_messages_provider_config",
|
||||
python_frame=r"ProviderConfigManager\.get_provider_anthropic_messages_config$",
|
||||
|
|
@ -30,10 +42,42 @@ COMMON_MAPPINGS: Final = (
|
|||
),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(rust_span="transform_response", python_frame=r"(?<!async_)transform_anthropic_messages_response$"),
|
||||
mapping(span="python_logging_pre_call", python_frame=r"Logging\.pre_call$"),
|
||||
mapping(span="python_logging_post_call", python_frame=r"Logging\.post_call$"),
|
||||
)
|
||||
|
||||
SUCCESS_MAPPINGS: Final = (mapping(span="python_success_callback", python_frame=r"Logging\.async_success_handler$"),)
|
||||
FAILURE_MAPPINGS: Final = (
|
||||
mapping(span="python_failure_callback", python_frame=r"Logging\.failure_handler$"),
|
||||
mapping(span="python_async_failure_callback", python_frame=r"Logging\.async_failure_handler$"),
|
||||
mapping(span="python_exception_mapping", python_frame=r"(?<!_)exception_type$"),
|
||||
)
|
||||
STREAM_MAPPINGS: Final = (
|
||||
mapping(span="python_stream_wrapper", python_frame=r"AnthropicMessagesStreamingResponse\.__init__$"),
|
||||
mapping(span="python_stream_next", python_frame=r"AnthropicMessagesStreamingResponse\.__anext__$"),
|
||||
mapping(
|
||||
span="python_stream_iterator",
|
||||
python_frame=r"BaseAnthropicMessagesStreamingIterator\.get_async_streaming_response_iterator$",
|
||||
),
|
||||
mapping(span="python_stream_chunks", python_frame=r"PassThroughStreamingHandler\.chunk_processor$"),
|
||||
mapping(
|
||||
span="python_stream_logging",
|
||||
python_frame=r"PassThroughStreamingHandler\._route_streaming_logging_to_handler$",
|
||||
),
|
||||
)
|
||||
|
||||
ANTHROPIC_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS,
|
||||
*SUCCESS_MAPPINGS,
|
||||
mapping(
|
||||
rust_span="transform_request",
|
||||
python_frame=r"(?<!Azure)AnthropicMessagesConfig\.transform_anthropic_messages_request$",
|
||||
),
|
||||
)
|
||||
|
||||
ANTHROPIC_FAILURE_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS,
|
||||
*FAILURE_MAPPINGS,
|
||||
mapping(
|
||||
rust_span="transform_request",
|
||||
python_frame=r"(?<!Azure)AnthropicMessagesConfig\.transform_anthropic_messages_request$",
|
||||
|
|
@ -42,6 +86,7 @@ ANTHROPIC_MAPPINGS: Final = (
|
|||
|
||||
AZURE_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS,
|
||||
*SUCCESS_MAPPINGS,
|
||||
mapping(
|
||||
rust_span="transform_request",
|
||||
python_frame=r"AzureAnthropicMessagesConfig\.transform_anthropic_messages_request$",
|
||||
|
|
@ -52,31 +97,54 @@ AZURE_MAPPINGS: Final = (
|
|||
),
|
||||
)
|
||||
|
||||
BEDROCK_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS,
|
||||
*SUCCESS_MAPPINGS,
|
||||
mapping(
|
||||
rust_span="transform_request",
|
||||
python_frame=r"AmazonAnthropicClaudeMessagesConfig\.transform_anthropic_messages_request$",
|
||||
),
|
||||
mapping(
|
||||
span="python_anthropic_transform_request",
|
||||
python_frame=r"(?<!Azure)AnthropicMessagesConfig\.transform_anthropic_messages_request$",
|
||||
),
|
||||
mapping(
|
||||
span="python_bedrock_provider_config",
|
||||
python_frame=r"BedrockModelInfo\.get_bedrock_provider_config_for_messages_api$",
|
||||
),
|
||||
mapping(span="python_aws_signing", python_frame=r"sign_request_off_loop_if_aws$"),
|
||||
mapping(
|
||||
span="python_aws_sign_request",
|
||||
python_frame=r"AmazonAnthropicClaudeMessagesConfig\.sign_request$|BaseAWSLLM\._sign_request$",
|
||||
),
|
||||
)
|
||||
|
||||
RETRY_MAPPINGS: Final = (
|
||||
*BEDROCK_MAPPINGS,
|
||||
mapping(
|
||||
span="python_retry_request_transform",
|
||||
python_frame=r"transform_anthropic_messages_request_on_http_error$",
|
||||
),
|
||||
mapping(
|
||||
span="python_strip_invalid_thinking",
|
||||
python_frame=r"strip_thinking_blocks_from_anthropic_messages_request_dict$",
|
||||
),
|
||||
)
|
||||
|
||||
MOCK_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS,
|
||||
mapping(span="python_mock_response", python_frame=r"messages/utils\.py:\d+ mock_response$"),
|
||||
)
|
||||
|
||||
|
||||
def _fixture(engine: Engine, provider: str) -> RouteFixture:
|
||||
conversation: Final = {"messages": [{"role": "user", "content": "hello"}], "max_tokens": 16}
|
||||
response: Final = json.dumps(
|
||||
{
|
||||
"id": "msg_trace",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-5",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 2, "output_tokens": 3},
|
||||
}
|
||||
).encode()
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
"model": f"{provider}/claude-sonnet-5",
|
||||
**({"body": {**conversation, "model": "claude-sonnet-5"}} if engine == "rust" else conversation),
|
||||
},
|
||||
provider_responses=(
|
||||
RecordedHttpResponse.from_bytes(
|
||||
200, (HttpHeader(name="content-type", value="application/json"),), response
|
||||
),
|
||||
),
|
||||
provider_responses=(json_response(anthropic_response_body()),),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -88,11 +156,170 @@ def _azure_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
|||
return _fixture(engine, "azure_ai")
|
||||
|
||||
|
||||
def _bedrock_kwargs(engine: Engine) -> dict[str, object]:
|
||||
conversation: Final = {"messages": [{"role": "user", "content": "hello"}], "max_tokens": 16}
|
||||
return {
|
||||
"model": "bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
**(
|
||||
{"body": {**conversation, "model": "anthropic.claude-3-sonnet-20240229-v1:0"}}
|
||||
if engine == "rust"
|
||||
else conversation
|
||||
),
|
||||
"aws_access_key_id": "AKIAIOSFODNN7EXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
|
||||
"aws_region_name": "us-east-1",
|
||||
}
|
||||
|
||||
|
||||
def _bedrock_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
response_fixture: Final = _fixture(engine, "anthropic")
|
||||
return RouteFixture(kwargs=_bedrock_kwargs(engine), provider_responses=response_fixture.provider_responses)
|
||||
|
||||
|
||||
def _bedrock_retry_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
success_fixture: Final = _bedrock_fixture(engine, _base_url)
|
||||
messages: Final = [
|
||||
{"role": "user", "content": "hello"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "thinking", "thinking": "old reasoning", "signature": ""},
|
||||
{"type": "text", "text": "partial answer"},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "continue"},
|
||||
]
|
||||
kwargs: Final = {
|
||||
**_bedrock_kwargs(engine),
|
||||
**(
|
||||
{"body": {"messages": messages, "max_tokens": 16, "model": "anthropic.claude-3-sonnet-20240229-v1:0"}}
|
||||
if engine == "rust"
|
||||
else {"messages": messages}
|
||||
),
|
||||
}
|
||||
return success_fixture.derive(
|
||||
kwargs=kwargs,
|
||||
provider_responses=(
|
||||
json_response({"message": "messages.1.content.0: Invalid `signature` in `thinking` block"}, status=400),
|
||||
*success_fixture.provider_responses,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _mock_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _fixture(engine, "anthropic")
|
||||
return fixture.derive(kwargs={"mock_response": "hello from mock"}, provider_responses=())
|
||||
|
||||
|
||||
def _provider_error_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _fixture(engine, "anthropic")
|
||||
return fixture.derive(
|
||||
provider_responses=(
|
||||
json_response(
|
||||
{"type": "error", "error": {"type": "invalid_request_error", "message": "bad request"}},
|
||||
status=400,
|
||||
),
|
||||
),
|
||||
expected_failure=True,
|
||||
)
|
||||
|
||||
|
||||
def _sync_unsupported_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
if engine == "rust":
|
||||
return _anthropic_fixture(engine, base_url)
|
||||
fixture: Final = _fixture(engine, "anthropic")
|
||||
return fixture.derive(provider_responses=(), expected_failure=True)
|
||||
|
||||
|
||||
def _stream_fixture_for(engine: Engine, provider: str) -> RouteFixture:
|
||||
fixture: Final = _fixture(engine, provider)
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(sse_response(anthropic_stream_events()),),
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
def _stream_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
return _stream_fixture_for(engine, "anthropic")
|
||||
|
||||
|
||||
def _azure_stream_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
return _stream_fixture_for(engine, "azure_ai")
|
||||
|
||||
|
||||
def _bedrock_stream_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
fixture: Final = _bedrock_fixture(engine, base_url)
|
||||
events: Final = tuple(payload for _, payload in anthropic_stream_events())
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(aws_event_stream_response(events),),
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
def _bedrock_stream_error_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
fixture: Final = _bedrock_fixture(engine, base_url)
|
||||
start: Final = anthropic_stream_events(model="anthropic.claude-3-sonnet-20240229-v1:0")[0][1]
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(aws_event_stream_response((start, {"type": "message_stop"}), corrupt_last_frame=True),),
|
||||
expected_failure=True,
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
SPEC: Final = RouteSpec("messages", ("create", "acreate"), ("messages", "amessages"), _anthropic_fixture)
|
||||
TRACE_SUITE: Final = TraceSuite(
|
||||
route=SPEC,
|
||||
scenarios=(
|
||||
TraceScenario(name="anthropic", fixture=_anthropic_fixture, mappings=ANTHROPIC_MAPPINGS, modes=("async",)),
|
||||
TraceScenario(name="azure-ai", fixture=_azure_fixture, mappings=AZURE_MAPPINGS, modes=("async",)),
|
||||
TraceScenario(
|
||||
name="async-anthropic", fixture=_anthropic_fixture, mappings=ANTHROPIC_MAPPINGS, asynchronous=True
|
||||
),
|
||||
TraceScenario(name="async-azure-ai", fixture=_azure_fixture, mappings=AZURE_MAPPINGS, asynchronous=True),
|
||||
TraceScenario(name="async-bedrock", fixture=_bedrock_fixture, mappings=BEDROCK_MAPPINGS, asynchronous=True),
|
||||
TraceScenario(
|
||||
name="async-bedrock-invalid-thinking-retry",
|
||||
fixture=_bedrock_retry_fixture,
|
||||
mappings=RETRY_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(name="async-mock-response", fixture=_mock_fixture, mappings=MOCK_MAPPINGS, asynchronous=True),
|
||||
TraceScenario(
|
||||
name="async-anthropic-provider-error",
|
||||
fixture=_provider_error_fixture,
|
||||
mappings=ANTHROPIC_FAILURE_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-anthropic-stream",
|
||||
fixture=_stream_fixture,
|
||||
mappings=(*ANTHROPIC_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-azure-ai-stream",
|
||||
fixture=_azure_stream_fixture,
|
||||
mappings=(*AZURE_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-bedrock-event-stream",
|
||||
fixture=_bedrock_stream_fixture,
|
||||
mappings=(*BEDROCK_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-bedrock-event-stream-error",
|
||||
fixture=_bedrock_stream_error_fixture,
|
||||
mappings=(*BEDROCK_MAPPINGS, *STREAM_MAPPINGS, *FAILURE_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-unsupported",
|
||||
fixture=_sync_unsupported_fixture,
|
||||
mappings=ANTHROPIC_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Final, cast
|
||||
from typing import Final
|
||||
|
||||
from .....shared.parity.recorded_http import HttpHeader, RecordedHttpResponse
|
||||
from .....shared.tracing.steps import Engine, mapping
|
||||
|
|
@ -53,6 +53,23 @@ ASYNC_MAPPINGS: Final = (
|
|||
),
|
||||
)
|
||||
|
||||
PUBLIC_RUST_DISPATCH_MAPPINGS: Final = (
|
||||
mapping(span="public_sdk_entrypoint", python_frame=r"ocr/main\.py:\d+ a?ocr$"),
|
||||
mapping(span="public_request", python_frame=r"ocr/main\.py:\d+ _public_request$"),
|
||||
mapping(span="bind_request", python_frame=r"ocr/main\.py:\d+ _bind_request$"),
|
||||
mapping(span="rust_ocr_enabled", python_frame=r"rust_bridge/configuration\.py:\d+ rust_ocr_enabled$"),
|
||||
mapping(span="select_native_ocr", python_frame=r"rust_bridge/ocr_lifecycle\.py:\d+ select$"),
|
||||
mapping(span="load_native_bridge", python_frame=r"rust_bridge/bindings\.py:\d+ NativeBinding\.load$"),
|
||||
mapping(span="native_call_setup", python_frame=r"rust_bridge/lifecycle\.py:\d+ setup$"),
|
||||
mapping(span="native_response", python_frame=r"rust_bridge/ocr\.py:\d+ _response$"),
|
||||
mapping(span="native_call_finalize", python_frame=r"rust_bridge/lifecycle\.py:\d+ finalize$"),
|
||||
mapping(
|
||||
span="native_success_bookkeeping",
|
||||
python_frame=r"rust_bridge/lifecycle\.py:\d+ success_bookkeeping$",
|
||||
),
|
||||
*(mapping(rust_span=item.rust) for item in SYNC_MAPPINGS if item.rust is not None),
|
||||
)
|
||||
|
||||
CALLBACK_SUCCESS_SYNC_MAPPINGS: Final = (*SYNC_MAPPINGS, SUCCESS_CALLBACK_SYNC_MAPPING)
|
||||
CALLBACK_SUCCESS_ASYNC_MAPPINGS: Final = (*ASYNC_MAPPINGS, SUCCESS_CALLBACK_ASYNC_MAPPING)
|
||||
CALLBACK_FAILURE_SYNC_MAPPINGS: Final = (
|
||||
|
|
@ -164,23 +181,6 @@ def _azure_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
|||
)
|
||||
|
||||
|
||||
def _vertex_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _fixture(
|
||||
engine,
|
||||
"vertex_ai/mistral-ocr-maas",
|
||||
{"type": "image_url", "image_url": "data:image/png;base64,aGVsbG8="},
|
||||
)
|
||||
vertex: Final = {"vertex_project": "trace-project", "vertex_location": "us-central1"}
|
||||
optional_params: Final = cast(dict[str, object], fixture.kwargs.get("optional_params", {}))
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
**fixture.kwargs,
|
||||
**({"optional_params": {**optional_params, **vertex}} if engine == "rust" else vertex),
|
||||
},
|
||||
provider_responses=fixture.provider_responses,
|
||||
)
|
||||
|
||||
|
||||
def _vertex_deepseek_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
vertex: Final = {"vertex_project": "trace-project", "vertex_location": "us-central1"}
|
||||
return RouteFixture(
|
||||
|
|
@ -204,6 +204,62 @@ def _vertex_deepseek_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
|||
)
|
||||
|
||||
|
||||
def _vertex_deepseek_credentials_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
|
||||
fixture: Final = _vertex_deepseek_fixture(engine, base_url)
|
||||
private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
credentials: Final = json.dumps(
|
||||
{
|
||||
"type": "service_account",
|
||||
"project_id": "trace-project",
|
||||
"private_key_id": "trace-key",
|
||||
"private_key": private_key.private_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PrivateFormat.PKCS8,
|
||||
serialization.NoEncryption(),
|
||||
).decode(),
|
||||
"client_email": "trace@trace-project.iam.gserviceaccount.com",
|
||||
"token_uri": f"{base_url}/token",
|
||||
}
|
||||
)
|
||||
return RouteFixture(
|
||||
kwargs={**fixture.kwargs, "api_key": None},
|
||||
environment=(("VERTEXAI_CREDENTIALS", credentials), ("VERTEX_AI_API_KEY", "")),
|
||||
provider_responses=(
|
||||
RecordedHttpResponse.from_bytes(
|
||||
200,
|
||||
(HttpHeader(name="content-type", value="application/json"),),
|
||||
b'{"access_token":"trace-token","token_type":"Bearer","expires_in":3600}',
|
||||
),
|
||||
*fixture.provider_responses,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _cohere_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
"model": "cohere/parse-v5.0",
|
||||
"document": {"type": "image_url", "image_url": "data:image/png;base64,aGVsbG8="},
|
||||
**({"optional_params": {"output_format": "blocks"}} if engine == "rust" else {"output_format": "blocks"}),
|
||||
},
|
||||
provider_responses=(
|
||||
RecordedHttpResponse.from_bytes(
|
||||
200,
|
||||
(HttpHeader(name="content-type", value="application/json"),),
|
||||
json.dumps(
|
||||
{
|
||||
"pages": [{"index": 0, "blocks": [{"type": "text", "text": {"content": "hello"}}]}],
|
||||
"meta": {"billed_units": {"pages": 1}},
|
||||
}
|
||||
).encode(),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _azure_document_intelligence_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
completed: Final = json.dumps(
|
||||
{
|
||||
|
|
@ -249,30 +305,6 @@ def _azure_document_intelligence_fixture(engine: Engine, base_url: str) -> Route
|
|||
)
|
||||
|
||||
|
||||
VERTEX_COMMON_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS[:7],
|
||||
mapping(
|
||||
rust_span="transform_ocr_request",
|
||||
python_frame=(
|
||||
r"VertexAIOCRConfig\.(?:async_)?transform_ocr_request$"
|
||||
r"|MistralOCRConfig\.transform_ocr_request$"
|
||||
),
|
||||
),
|
||||
COMMON_MAPPINGS[-1],
|
||||
)
|
||||
VERTEX_SYNC_MAPPINGS: Final = (
|
||||
*VERTEX_COMMON_MAPPINGS,
|
||||
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
|
||||
mapping(span="python_transform_ocr_response_wrapper", python_frame=r"BaseLLMHTTPHandler\._transform_ocr_response$"),
|
||||
mapping(rust_span="transform_ocr_response", python_frame=r"MistralOCRConfig\.transform_ocr_response$"),
|
||||
)
|
||||
VERTEX_ASYNC_MAPPINGS: Final = (
|
||||
*VERTEX_COMMON_MAPPINGS,
|
||||
mapping(span="python_ocr_wrapper", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
|
||||
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.async_ocr$"),
|
||||
mapping(rust_span="transform_ocr_response", python_frame=r"MistralOCRConfig\.transform_ocr_response$"),
|
||||
)
|
||||
|
||||
DEEPSEEK_COMMON_MAPPINGS: Final = (
|
||||
mapping(rust_span="ocr", python_frame=r"ocr/main\.py:\d+ a?ocr$"),
|
||||
mapping(rust_span="prepare_ocr_call", python_frame=r"ocr/main\.py:\d+ _prepare_ocr_request$"),
|
||||
|
|
@ -352,59 +384,126 @@ DOCUMENT_INTELLIGENCE_ASYNC_MAPPINGS: Final = (
|
|||
mapping(span="python_poll_http_request", python_frame=r"AsyncHTTPHandler\.get$"),
|
||||
)
|
||||
|
||||
COHERE_COMMON_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS[:7],
|
||||
mapping(
|
||||
rust_span="transform_ocr_request",
|
||||
python_frame=r"CohereParseConfig\.(?:async_)?transform_ocr_request$",
|
||||
),
|
||||
COMMON_MAPPINGS[-1],
|
||||
mapping(rust_span="transform_ocr_response", python_frame=r"CohereParseConfig\.transform_ocr_response$"),
|
||||
)
|
||||
COHERE_ASYNC_MAPPINGS: Final = (
|
||||
*COHERE_COMMON_MAPPINGS,
|
||||
mapping(span="python_ocr_wrapper", python_frame=r"BaseLLMHTTPHandler\.ocr$"),
|
||||
mapping(rust_span="execute_ocr_provider_call", python_frame=r"BaseLLMHTTPHandler\.async_ocr$"),
|
||||
)
|
||||
|
||||
SPEC: Final = RouteSpec("ocr", ("ocr", "aocr"), ("ocr", "aocr"), _mistral_fixture)
|
||||
TRACE_SUITE: Final = TraceSuite(
|
||||
route=SPEC,
|
||||
scenarios=(
|
||||
TraceScenario(
|
||||
name="mistral",
|
||||
name="sync-mistral",
|
||||
fixture=_mistral_fixture,
|
||||
mappings=COMMON_MAPPINGS,
|
||||
sync_mappings=(*SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
async_mappings=(*ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
mappings=(*SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="mistral-callback-success",
|
||||
name="async-mistral",
|
||||
fixture=_mistral_fixture,
|
||||
mappings=(*ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-mistral-callback-success",
|
||||
fixture=_mistral_callback_success_fixture,
|
||||
mappings=COMMON_MAPPINGS,
|
||||
sync_mappings=CALLBACK_SUCCESS_SYNC_MAPPINGS,
|
||||
async_mappings=CALLBACK_SUCCESS_ASYNC_MAPPINGS,
|
||||
mappings=CALLBACK_SUCCESS_SYNC_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="mistral-callback-failure",
|
||||
name="async-mistral-callback-success",
|
||||
fixture=_mistral_callback_success_fixture,
|
||||
mappings=CALLBACK_SUCCESS_ASYNC_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-mistral-callback-failure",
|
||||
fixture=_mistral_callback_failure_fixture,
|
||||
mappings=(*COMMON_MAPPINGS, FAILURE_CALLBACK_MAPPING),
|
||||
sync_mappings=CALLBACK_FAILURE_SYNC_MAPPINGS,
|
||||
async_mappings=CALLBACK_FAILURE_ASYNC_MAPPINGS,
|
||||
mappings=CALLBACK_FAILURE_SYNC_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="azure-ai",
|
||||
name="async-mistral-callback-failure",
|
||||
fixture=_mistral_callback_failure_fixture,
|
||||
mappings=CALLBACK_FAILURE_ASYNC_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-azure-ai",
|
||||
fixture=_azure_fixture,
|
||||
mappings=AZURE_COMMON_MAPPINGS,
|
||||
sync_mappings=(*AZURE_SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
async_mappings=(*AZURE_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
mappings=(*AZURE_SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="azure-document-intelligence",
|
||||
name="async-azure-ai",
|
||||
fixture=_azure_fixture,
|
||||
mappings=(*AZURE_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-azure-document-intelligence",
|
||||
fixture=_azure_document_intelligence_fixture,
|
||||
mappings=DOCUMENT_INTELLIGENCE_COMMON_MAPPINGS,
|
||||
sync_mappings=(*DOCUMENT_INTELLIGENCE_SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
async_mappings=(*DOCUMENT_INTELLIGENCE_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
mappings=(*DOCUMENT_INTELLIGENCE_SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="vertex-ai",
|
||||
fixture=_vertex_fixture,
|
||||
mappings=VERTEX_COMMON_MAPPINGS,
|
||||
sync_mappings=(*VERTEX_SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
async_mappings=(*VERTEX_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
name="async-azure-document-intelligence",
|
||||
fixture=_azure_document_intelligence_fixture,
|
||||
mappings=(*DOCUMENT_INTELLIGENCE_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="vertex-deepseek",
|
||||
name="sync-vertex-deepseek",
|
||||
fixture=_vertex_deepseek_fixture,
|
||||
mappings=DEEPSEEK_COMMON_MAPPINGS,
|
||||
sync_mappings=(*DEEPSEEK_SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
async_mappings=(*DEEPSEEK_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
mappings=(*DEEPSEEK_SYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-vertex-deepseek",
|
||||
fixture=_vertex_deepseek_fixture,
|
||||
mappings=(*DEEPSEEK_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-vertex-deepseek-credentials",
|
||||
fixture=_vertex_deepseek_credentials_fixture,
|
||||
mappings=DEEPSEEK_SYNC_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-vertex-deepseek-credentials",
|
||||
fixture=_vertex_deepseek_credentials_fixture,
|
||||
mappings=DEEPSEEK_ASYNC_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-cohere",
|
||||
fixture=_cohere_fixture,
|
||||
mappings=(*COHERE_ASYNC_MAPPINGS, IGNORED_SUCCESS_CALLBACK_MAPPING),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="sync-public-rust-dispatch",
|
||||
fixture=_mistral_fixture,
|
||||
mappings=PUBLIC_RUST_DISPATCH_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-public-rust-dispatch",
|
||||
fixture=_mistral_fixture,
|
||||
mappings=PUBLIC_RUST_DISPATCH_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,216 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
from .....shared.tracing.steps import Engine, mapping
|
||||
from ...fixtures import (
|
||||
anthropic_response_body,
|
||||
anthropic_stream_events,
|
||||
json_response,
|
||||
responses_body,
|
||||
responses_stream_events,
|
||||
sse_response,
|
||||
)
|
||||
from ...models import RouteFixture, RouteSpec, TraceScenario, TraceSuite
|
||||
|
||||
COMMON_MAPPINGS: Final = (
|
||||
mapping(span="python_responses", python_frame=r"responses/main\.py:\d+ a?responses$"),
|
||||
mapping(
|
||||
span="python_responses_provider_config",
|
||||
python_frame=r"ProviderConfigManager\.get_provider_responses_api_config$",
|
||||
),
|
||||
mapping(rust_span="responses_provider_config"),
|
||||
mapping(rust_span="validate_environment", python_frame=r"validate_environment$"),
|
||||
mapping(rust_span="complete_url", python_frame=r"get_complete_url$"),
|
||||
mapping(
|
||||
rust_span="transform_request",
|
||||
python_frame=r"(?<!AzureOpenAIResponsesAPIConfig\.)transform_responses_api_request$",
|
||||
),
|
||||
mapping(
|
||||
rust_span="execute_responses_provider_call",
|
||||
python_frame=r"BaseLLMHTTPHandler\.(?:async_)?response_api_handler$",
|
||||
),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(rust_span="transform_response", python_frame=r"transform_response_api_response$"),
|
||||
mapping(span="python_logging_pre_call", python_frame=r"Logging\.pre_call$"),
|
||||
mapping(span="python_success_callback", python_frame=r"Logging\.async_success_handler$|Logging\.success_handler$"),
|
||||
)
|
||||
|
||||
STREAM_MAPPINGS: Final = (
|
||||
mapping(
|
||||
span="python_responses_stream_iterator",
|
||||
python_frame=r"(?:Sync)?ResponsesAPIStreamingIterator\.__init__$",
|
||||
),
|
||||
mapping(
|
||||
span="python_responses_stream_next",
|
||||
python_frame=r"(?:Sync)?ResponsesAPIStreamingIterator\.__a?next__$",
|
||||
),
|
||||
mapping(span="python_responses_stream_transform", python_frame=r"transform_streaming_response$"),
|
||||
)
|
||||
|
||||
FAILURE_MAPPINGS: Final = (
|
||||
mapping(span="python_exception_mapping", python_frame=r"(?<!_)exception_type$"),
|
||||
mapping(span="python_failure_callback", python_frame=r"Logging\.failure_handler$"),
|
||||
mapping(span="python_async_failure_callback", python_frame=r"Logging\.async_failure_handler$"),
|
||||
)
|
||||
|
||||
AZURE_MAPPINGS: Final = (
|
||||
*COMMON_MAPPINGS,
|
||||
mapping(
|
||||
span="python_azure_transform_request",
|
||||
python_frame=r"AzureOpenAIResponsesAPIConfig\.transform_responses_api_request$",
|
||||
),
|
||||
)
|
||||
|
||||
BRIDGE_MAPPINGS: Final = (
|
||||
mapping(span="python_responses", python_frame=r"responses/main\.py:\d+ a?responses$"),
|
||||
mapping(
|
||||
span="python_responses_chat_bridge", python_frame=r"ResponsesToCompletionBridgeHandler\.response_api_handler$"
|
||||
),
|
||||
mapping(span="python_chat_completions", python_frame=r"main\.py:\d+ a?completion$"),
|
||||
mapping(span="python_chat_transform_request", python_frame=r"AnthropicConfig\.transform_request$"),
|
||||
mapping(span="python_logging_pre_call", python_frame=r"Logging\.pre_call$"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(span="python_chat_transform_response", python_frame=r"AnthropicConfig\.transform_response$"),
|
||||
mapping(span="python_chat_to_responses", python_frame=r"LiteLLMResponsesTransformationHandler\..*response"),
|
||||
mapping(span="python_success_callback", python_frame=r"Logging\.async_success_handler$|Logging\.success_handler$"),
|
||||
)
|
||||
|
||||
|
||||
def _native_fixture(engine: Engine, provider: str) -> RouteFixture:
|
||||
model: Final = "gpt-5"
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
"model": f"{provider}/{model}",
|
||||
"input": "hello",
|
||||
**({"body": {"model": model, "input": "hello"}} if engine == "rust" else {}),
|
||||
},
|
||||
provider_responses=(json_response(responses_body(model=model)),),
|
||||
)
|
||||
|
||||
|
||||
def _openai_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
return _native_fixture(engine, "openai")
|
||||
|
||||
|
||||
def _azure_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _native_fixture(engine, "azure")
|
||||
return fixture.derive(kwargs={"api_version": "2025-04-01-preview"})
|
||||
|
||||
|
||||
def _openai_stream_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _openai_fixture(engine, _base_url)
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(sse_response(responses_stream_events()),),
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
def _provider_error_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _openai_fixture(engine, _base_url)
|
||||
return fixture.derive(
|
||||
provider_responses=(
|
||||
json_response({"error": {"message": "bad request", "type": "invalid_request_error"}}, status=400),
|
||||
),
|
||||
expected_failure=True,
|
||||
)
|
||||
|
||||
|
||||
def _stream_failed_fixture(engine: Engine, base_url: str) -> RouteFixture:
|
||||
fixture: Final = _openai_fixture(engine, base_url)
|
||||
failed_response: Final[dict[str, object]] = {
|
||||
**responses_body(),
|
||||
"status": "failed",
|
||||
"output": [],
|
||||
"error": {"message": "stream failed", "type": "server_error", "code": "server_error"},
|
||||
}
|
||||
events: Final = (
|
||||
(
|
||||
"response.created",
|
||||
{"type": "response.created", "response": {**failed_response, "status": "in_progress", "error": None}},
|
||||
),
|
||||
("response.failed", {"type": "response.failed", "response": failed_response}),
|
||||
)
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(sse_response(events),),
|
||||
expected_failure=True,
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_bridge_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
return RouteFixture(
|
||||
kwargs={
|
||||
"model": "anthropic/claude-sonnet-5",
|
||||
"input": "hello",
|
||||
"max_output_tokens": 16,
|
||||
**({"body": {"model": "claude-sonnet-5", "input": "hello"}} if engine == "rust" else {}),
|
||||
},
|
||||
provider_responses=(json_response(anthropic_response_body()),),
|
||||
)
|
||||
|
||||
|
||||
def _anthropic_bridge_stream_fixture(engine: Engine, _base_url: str) -> RouteFixture:
|
||||
fixture: Final = _anthropic_bridge_fixture(engine, _base_url)
|
||||
return fixture.derive(
|
||||
kwargs={"stream": True},
|
||||
provider_responses=(sse_response(anthropic_stream_events()),),
|
||||
consume_stream=True,
|
||||
)
|
||||
|
||||
|
||||
SPEC: Final = RouteSpec("responses", ("responses", "aresponses"), None, _openai_fixture)
|
||||
TRACE_SUITE: Final = TraceSuite(
|
||||
route=SPEC,
|
||||
scenarios=(
|
||||
TraceScenario(name="sync-openai", fixture=_openai_fixture, mappings=COMMON_MAPPINGS, asynchronous=False),
|
||||
TraceScenario(name="async-openai", fixture=_openai_fixture, mappings=COMMON_MAPPINGS, asynchronous=True),
|
||||
TraceScenario(
|
||||
name="sync-openai-stream",
|
||||
fixture=_openai_stream_fixture,
|
||||
mappings=(*COMMON_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-openai-stream",
|
||||
fixture=_openai_stream_fixture,
|
||||
mappings=(*COMMON_MAPPINGS, *STREAM_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-openai-provider-error",
|
||||
fixture=_provider_error_fixture,
|
||||
mappings=(*COMMON_MAPPINGS, *FAILURE_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-openai-stream-failed",
|
||||
fixture=_stream_failed_fixture,
|
||||
mappings=(*COMMON_MAPPINGS, *STREAM_MAPPINGS, *FAILURE_MAPPINGS),
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(name="async-azure", fixture=_azure_fixture, mappings=AZURE_MAPPINGS, asynchronous=True),
|
||||
TraceScenario(
|
||||
name="async-anthropic-chat-bridge",
|
||||
fixture=_anthropic_bridge_fixture,
|
||||
mappings=BRIDGE_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-anthropic-chat-bridge-stream",
|
||||
fixture=_anthropic_bridge_stream_fixture,
|
||||
mappings=(
|
||||
*BRIDGE_MAPPINGS,
|
||||
mapping(span="python_chat_stream_wrapper", python_frame=r"CustomStreamWrapper\.__init__$"),
|
||||
mapping(span="python_chat_stream_next", python_frame=r"CustomStreamWrapper\.__anext__$"),
|
||||
mapping(
|
||||
span="python_responses_bridge_stream_iterator",
|
||||
python_frame=r"LiteLLMCompletionStreamingIterator\.__init__$|LiteLLMCompletionStreamingIterator\.__anext__$",
|
||||
),
|
||||
),
|
||||
asynchronous=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
@ -0,0 +1,69 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import Final, cast
|
||||
|
||||
from ..models import TraceSuite
|
||||
|
||||
|
||||
def _suite(module: str) -> TraceSuite:
|
||||
loaded: Final = import_module(module)
|
||||
candidate: Final = cast(object, getattr(loaded, "TRACE_SUITE"))
|
||||
assert isinstance(candidate, TraceSuite)
|
||||
return candidate
|
||||
|
||||
|
||||
def test_core_sdk_scenario_matrix_keeps_distinct_migration_paths() -> None:
|
||||
chat: Final = _suite("tests.rust-python-harness.strategies.trace_parity.sdk.chat_completions.case")
|
||||
messages: Final = _suite("tests.rust-python-harness.strategies.trace_parity.sdk.messages.case")
|
||||
ocr: Final = _suite("tests.rust-python-harness.strategies.trace_parity.sdk.ocr.case")
|
||||
responses: Final = _suite("tests.rust-python-harness.strategies.trace_parity.sdk.responses.case")
|
||||
|
||||
assert {(scenario.name, scenario.asynchronous) for scenario in chat.scenarios} >= {
|
||||
("sync-anthropic", False),
|
||||
("async-anthropic", True),
|
||||
("sync-anthropic-stream", False),
|
||||
("async-anthropic-stream", True),
|
||||
("async-anthropic-provider-error", True),
|
||||
("async-anthropic-stream-error", True),
|
||||
("sync-bedrock", False),
|
||||
("async-bedrock", True),
|
||||
("sync-bedrock-event-stream", False),
|
||||
("async-bedrock-event-stream", True),
|
||||
}
|
||||
assert {(scenario.name, scenario.asynchronous) for scenario in messages.scenarios} >= {
|
||||
("async-anthropic-stream", True),
|
||||
("async-azure-ai-stream", True),
|
||||
("async-bedrock-event-stream", True),
|
||||
("async-bedrock-event-stream-error", True),
|
||||
("async-bedrock-invalid-thinking-retry", True),
|
||||
("sync-unsupported", False),
|
||||
}
|
||||
assert {(scenario.name, scenario.asynchronous) for scenario in ocr.scenarios} >= {
|
||||
("async-cohere", True),
|
||||
("sync-public-rust-dispatch", False),
|
||||
("async-public-rust-dispatch", True),
|
||||
}
|
||||
assert {(scenario.name, scenario.asynchronous) for scenario in responses.scenarios} >= {
|
||||
("sync-openai", False),
|
||||
("async-openai", True),
|
||||
("sync-openai-stream", False),
|
||||
("async-openai-stream", True),
|
||||
("async-openai-provider-error", True),
|
||||
("async-openai-stream-failed", True),
|
||||
("async-azure", True),
|
||||
("async-anthropic-chat-bridge", True),
|
||||
("async-anthropic-chat-bridge-stream", True),
|
||||
}
|
||||
|
||||
|
||||
def test_core_gateway_matrix_keeps_downstream_streams_separate() -> None:
|
||||
modules: Final = (
|
||||
"tests.rust-python-harness.strategies.trace_parity.gateway.chat_completions.case",
|
||||
"tests.rust-python-harness.strategies.trace_parity.gateway.messages.case",
|
||||
"tests.rust-python-harness.strategies.trace_parity.gateway.responses.case",
|
||||
)
|
||||
|
||||
for module in modules:
|
||||
suite = _suite(module)
|
||||
assert any("downstream-stream" in scenario.name for scenario in suite.scenarios)
|
||||
|
|
@ -92,11 +92,16 @@ TRACE_SUITE: Final = TraceSuite(
|
|||
route=SPEC,
|
||||
scenarios=(
|
||||
TraceScenario(
|
||||
name="bedrock",
|
||||
name="sync-bedrock",
|
||||
fixture=_fixture,
|
||||
mappings=MAPPINGS,
|
||||
sync_mappings=SYNC_MAPPINGS,
|
||||
async_mappings=ASYNC_MAPPINGS,
|
||||
mappings=SYNC_MAPPINGS,
|
||||
asynchronous=False,
|
||||
),
|
||||
TraceScenario(
|
||||
name="async-bedrock",
|
||||
fixture=_fixture,
|
||||
mappings=ASYNC_MAPPINGS,
|
||||
asynchronous=True,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,58 +1,46 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
|
||||
from ...shared.reporting.models import CaseResult, Coverage, HarnessCase, ResultArtifact, RunStatus
|
||||
from ...shared.reporting.strategy import ModuleCaseSpec, NotImplementedCaseSpec
|
||||
from ...shared.tracing.steps import PipelineStep, TraceContract, TraceMapping, mapping
|
||||
from ...shared.tracing.steps import PipelineStep
|
||||
from . import reporting
|
||||
from .reporting import TRACE_COMPARISON_ARTIFACT, TraceComparisonArtifact, render_trace_results
|
||||
|
||||
MAPPINGS: Final = (
|
||||
mapping(rust_span="ocr", python_frame=r"ocr/main\.py:\d+ a?ocr$"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$"),
|
||||
)
|
||||
from .reporting import TRACE_ARTIFACT, TraceArtifact, render_trace_results
|
||||
|
||||
|
||||
def _result(comparison: TraceComparisonArtifact) -> CaseResult:
|
||||
def _result(trace: TraceArtifact) -> CaseResult:
|
||||
case: Final = HarnessCase(
|
||||
strategy_id="trace_parity",
|
||||
strategy_label="Trace parity",
|
||||
sdk_function=comparison.sdk_function,
|
||||
sdk_function=trace.sdk_function,
|
||||
spec=ModuleCaseSpec(coverage=Coverage.PARTIAL, module="example"),
|
||||
surface=comparison.surface,
|
||||
surface=trace.surface,
|
||||
)
|
||||
result: Final = CaseResult(case=case)
|
||||
nodeid: Final = f"trace:sdk:{comparison.sdk_function}:{comparison.scenario}:{comparison.mode}"
|
||||
nodeid: Final = f"trace:{trace.surface}:{trace.sdk_function}:{trace.scenario}"
|
||||
result.collected.add(nodeid)
|
||||
result.record(
|
||||
nodeid,
|
||||
RunStatus.PASSED,
|
||||
artifacts=(ResultArtifact(TRACE_COMPARISON_ARTIFACT, comparison.model_dump_json()),),
|
||||
)
|
||||
result.record(nodeid, RunStatus.PASSED, artifacts=(ResultArtifact(TRACE_ARTIFACT, trace.model_dump_json()),))
|
||||
return result
|
||||
|
||||
|
||||
def _comparison(
|
||||
def _trace(
|
||||
python: tuple[PipelineStep, ...],
|
||||
rust: tuple[PipelineStep, ...],
|
||||
*,
|
||||
mappings: Sequence[TraceMapping] = MAPPINGS,
|
||||
rust_error: str | None = None,
|
||||
) -> TraceComparisonArtifact:
|
||||
return TraceComparisonArtifact.from_traces(
|
||||
engine: Literal["python", "rust", "both"] = "both",
|
||||
scenario: str = "sync-default",
|
||||
) -> TraceArtifact:
|
||||
return TraceArtifact.from_traces(
|
||||
engine=engine,
|
||||
surface="sdk",
|
||||
sdk_function="ocr",
|
||||
scenario="default",
|
||||
mode="sync",
|
||||
mappings=mappings,
|
||||
contract=TraceContract(),
|
||||
scenario=scenario,
|
||||
python=python,
|
||||
rust=rust,
|
||||
python_unmatched=796,
|
||||
rust_error=rust_error,
|
||||
)
|
||||
|
||||
|
|
@ -67,107 +55,55 @@ def _events(*items: tuple[str, int, str | None]) -> tuple[PipelineStep, ...]:
|
|||
return tuple(steps)
|
||||
|
||||
|
||||
def test_renderer_shows_matching_python_and_rust_paths() -> None:
|
||||
rust: Final = _events(("ocr", 0, None), ("http_request", 1, None))
|
||||
def test_renderer_prints_python_and_rust_traces_independently() -> None:
|
||||
python: Final = _events(
|
||||
("ocr", 0, "ocr/main.py:88 aocr"),
|
||||
("http_request", 1, "http_handler.py:673 AsyncHTTPHandler.post"),
|
||||
("python_prepare", 1, "prep.py:1 python_prepare"),
|
||||
)
|
||||
|
||||
section: Final = render_trace_results((_result(_comparison(python, rust)),))[0]
|
||||
report: Final = "\n\n".join(section.blocks)
|
||||
|
||||
assert section.title == "SDK trace comparisons"
|
||||
assert "Case: ocr" in report
|
||||
assert "PYTHON (2 steps)\n1 aocr (ocr/main.py:88)\n2 AsyncHTTPHandler.post (http_handler.py:673)" in report
|
||||
assert "RUST (2 steps)\nocr -> 1 aocr\n http_request -> 2 AsyncHTTPHandler.post" in report
|
||||
assert "Mapping (identifier -> span)" not in report
|
||||
assert "Trace: MATCH" in report
|
||||
assert "Same steps, order, and nesting" in report
|
||||
assert "Unseen mappings:" not in report
|
||||
|
||||
|
||||
def test_renderer_reports_mappings_that_matched_nothing() -> None:
|
||||
events: Final = _events(("ocr", 0, None))
|
||||
|
||||
section: Final = render_trace_results((_result(_comparison(events, events)),))[0]
|
||||
report: Final = "\n\n".join(section.blocks)
|
||||
|
||||
assert "Unseen mappings: http_request" in report
|
||||
assert "Contract: FAIL" in report
|
||||
|
||||
|
||||
def test_renderer_numbers_repeated_span_occurrences() -> None:
|
||||
mappings: Final = (MAPPINGS[0], MAPPINGS[1])
|
||||
rust: Final = _events(("ocr", 0, None), ("http_request", 1, None), ("http_request", 1, None))
|
||||
python: Final = _events(
|
||||
("ocr", 0, "ocr/main.py:88 aocr"),
|
||||
("http_request", 1, "http_handler.py:673 AsyncHTTPHandler.post"),
|
||||
("http_request", 1, "http_handler.py:673 AsyncHTTPHandler.post"),
|
||||
)
|
||||
|
||||
report: Final = "\n\n".join(render_trace_results((_result(_comparison(python, rust, mappings=mappings)),))[0].blocks)
|
||||
|
||||
assert "http_request#2" in report
|
||||
|
||||
|
||||
def test_renderer_accepts_declared_engine_specific_steps() -> None:
|
||||
mappings: Final = (
|
||||
*MAPPINGS[:1],
|
||||
mapping(span="python_prepare", python_frame=r"python_prepare$"),
|
||||
mapping(rust_span="rust_prepare"),
|
||||
)
|
||||
python: Final = _events(("ocr", 0, None), ("python_prepare", 1, "prep.py:1 python_prepare"))
|
||||
rust: Final = _events(("ocr", 0, None), ("rust_prepare", 1, None))
|
||||
|
||||
section: Final = render_trace_results((_result(_comparison(python, rust, mappings=mappings)),))[0]
|
||||
section: Final = render_trace_results((_result(_trace(python, rust)),))[0]
|
||||
report: Final = "\n\n".join(section.blocks)
|
||||
|
||||
assert "2 python_prepare (prep.py:1) [python only]" in report
|
||||
assert "rust_prepare -> [rust only]" in report
|
||||
assert "Trace: MATCH" in report
|
||||
assert "Contract: PASS" in report
|
||||
assert section.title == "SDK traces"
|
||||
assert "PYTHON (2 steps)\n1 aocr (ocr/main.py:88)\n2 python_prepare (prep.py:1)" in report
|
||||
assert "RUST (2 steps)\n1 ocr\n2 rust_prepare" in report
|
||||
assert "python only" not in report
|
||||
assert "rust only" not in report
|
||||
assert " -> " not in report
|
||||
assert "Trace: MATCH" not in report
|
||||
assert "Trace: DRIFT" not in report
|
||||
assert "Contract:" not in report
|
||||
|
||||
|
||||
def test_unavailable_check_reports_mode_from_nodeid() -> None:
|
||||
case: Final = HarnessCase(
|
||||
strategy_id="trace_parity",
|
||||
strategy_label="Trace parity",
|
||||
sdk_function="ocr",
|
||||
spec=ModuleCaseSpec(coverage=Coverage.PARTIAL, module="example"),
|
||||
surface="sdk",
|
||||
)
|
||||
result: Final = CaseResult(case=case)
|
||||
result.collected.add("trace:sdk:ocr:default:sync")
|
||||
result.record("trace:sdk:ocr:default:sync", RunStatus.ERROR)
|
||||
@pytest.mark.parametrize(
|
||||
("engine", "present", "absent"),
|
||||
(("python", "PYTHON (1 steps)", "RUST"), ("rust", "RUST (1 steps)", "PYTHON")),
|
||||
)
|
||||
def test_renderer_prints_only_selected_engine(engine: Literal["python", "rust"], present: str, absent: str) -> None:
|
||||
events: Final = _events(("ocr", 0, None))
|
||||
|
||||
section: Final = render_trace_results((result,))[0]
|
||||
report: Final = "\n\n".join(section.blocks)
|
||||
report: Final = "\n\n".join(render_trace_results((_result(_trace(events, events, engine=engine)),))[0].blocks)
|
||||
|
||||
assert "Case: ocr" in report
|
||||
assert "Scenario: default / Mode: sync" in report
|
||||
assert "Trace: NOT AVAILABLE\nTest outcome: error" in report
|
||||
assert "unknown mode" not in report
|
||||
assert present in report
|
||||
assert absent not in report
|
||||
|
||||
|
||||
def test_renderer_keeps_collected_trace_when_one_engine_errors() -> None:
|
||||
python: Final = _events(
|
||||
("ocr", 0, "ocr/main.py:88 aocr"),
|
||||
("http_request", 1, "http_handler.py:673 AsyncHTTPHandler.post"),
|
||||
python: Final = _events(("ocr", 0, "ocr/main.py:88 aocr"))
|
||||
|
||||
report: Final = "\n\n".join(
|
||||
render_trace_results(
|
||||
(_result(_trace(python, (), rust_error="rust: native Rust bridge must include the trace-parity feature")),)
|
||||
)[0].blocks
|
||||
)
|
||||
|
||||
section: Final = render_trace_results(
|
||||
(_result(_comparison(python, (), rust_error="rust: native Rust bridge must include the trace-parity feature")),)
|
||||
)[0]
|
||||
report: Final = "\n\n".join(section.blocks)
|
||||
|
||||
assert "PYTHON (2 steps)\n1 aocr (ocr/main.py:88) [python only]" in report
|
||||
assert "PYTHON (1 steps)\n1 aocr (ocr/main.py:88)" in report
|
||||
assert "Rust error: rust: native Rust bridge must include the trace-parity feature" in report
|
||||
assert "hint: rebuild the native bridge with the trace-parity feature" in report
|
||||
assert "Contract: FAIL" in report
|
||||
|
||||
|
||||
def test_renderer_groups_all_modes_under_one_case_header() -> None:
|
||||
def test_unavailable_trace_reports_scenario_from_nodeid() -> None:
|
||||
case: Final = HarnessCase(
|
||||
strategy_id="trace_parity",
|
||||
strategy_label="Trace parity",
|
||||
|
|
@ -176,76 +112,55 @@ def test_renderer_groups_all_modes_under_one_case_header() -> None:
|
|||
surface="sdk",
|
||||
)
|
||||
result: Final = CaseResult(case=case)
|
||||
events: Final = _events(("ocr", 0, None))
|
||||
modes: Final[tuple[Literal["sync", "async"], ...]] = ("sync", "async")
|
||||
for mode in modes:
|
||||
nodeid = f"trace:sdk:ocr:default:{mode}"
|
||||
result.collected.add(nodeid)
|
||||
comparison = TraceComparisonArtifact.from_traces(
|
||||
surface="sdk",
|
||||
sdk_function="ocr",
|
||||
scenario="default",
|
||||
mode=mode,
|
||||
mappings=MAPPINGS,
|
||||
contract=TraceContract(),
|
||||
python=events,
|
||||
rust=events,
|
||||
python_unmatched=0,
|
||||
)
|
||||
result.record(
|
||||
nodeid,
|
||||
RunStatus.PASSED,
|
||||
artifacts=(ResultArtifact(TRACE_COMPARISON_ARTIFACT, comparison.model_dump_json()),),
|
||||
)
|
||||
result.collected.add("trace:sdk:ocr:async-error")
|
||||
result.record("trace:sdk:ocr:async-error", RunStatus.ERROR)
|
||||
|
||||
section: Final = render_trace_results((result,))[0]
|
||||
report: Final = "\n\n".join(render_trace_results((result,))[0].blocks)
|
||||
|
||||
assert "Scenario: async-error" in report
|
||||
assert "Trace: NOT AVAILABLE\nTest outcome: error" in report
|
||||
|
||||
|
||||
def test_renderer_groups_scenarios_under_one_case_header() -> None:
|
||||
result: Final = _result(_trace(_events(("ocr", 0, None)), (), scenario="sync-default"))
|
||||
async_trace: Final = _trace((), _events(("ocr", 0, None)), scenario="async-default")
|
||||
nodeid: Final = "trace:sdk:ocr:async-default"
|
||||
result.collected.add(nodeid)
|
||||
result.record(nodeid, RunStatus.PASSED, artifacts=(ResultArtifact(TRACE_ARTIFACT, async_trace.model_dump_json()),))
|
||||
|
||||
report: Final = render_trace_results((result,))[0].blocks[0]
|
||||
|
||||
assert len(section.blocks) == 1
|
||||
report: Final = section.blocks[0]
|
||||
assert report.count("Case: ocr") == 1
|
||||
assert "Scenario: default / Mode: sync" in report
|
||||
assert "Scenario: default / Mode: async" in report
|
||||
assert "Scenario: sync-default" in report
|
||||
assert "Scenario: async-default" in report
|
||||
|
||||
|
||||
def test_renderer_colors_every_trace_line_in_a_terminal(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
rust: Final = _events(("ocr", 0, None), ("http_request", 1, None))
|
||||
python: Final = _events(
|
||||
("ocr", 0, "ocr/main.py:88 aocr"),
|
||||
("http_request", 1, "http_handler.py:673 AsyncHTTPHandler.post"),
|
||||
)
|
||||
events: Final = _events(("ocr", 0, "ocr/main.py:88 aocr"))
|
||||
monkeypatch.setattr(reporting.sys.stdout, "isatty", lambda: True)
|
||||
monkeypatch.delenv("NO_COLOR", raising=False)
|
||||
|
||||
section: Final = render_trace_results((_result(_comparison(python, rust)),))[0]
|
||||
report: Final = "\n\n".join(section.blocks)
|
||||
report: Final = "\n\n".join(render_trace_results((_result(_trace(events, events)),))[0].blocks)
|
||||
|
||||
assert "\033[36mPYTHON\033[0m (2 steps)" in report
|
||||
assert "\033[36mPYTHON\033[0m (1 steps)" in report
|
||||
assert "\033[36m1 aocr (ocr/main.py:88)\033[0m" in report
|
||||
assert "\033[33mRUST\033[0m (2 steps)" in report
|
||||
assert "\033[33mocr\033[0m -> \033[36m1 aocr\033[0m" in report
|
||||
assert "\033[33mhttp_request\033[0m -> \033[36m2 AsyncHTTPHandler.post\033[0m" in report
|
||||
assert "\033[33mRUST\033[0m (1 steps)" in report
|
||||
assert "\033[33m1 ocr\033[0m" in report
|
||||
|
||||
|
||||
def test_renderer_groups_cases_and_unavailable_entries_by_surface() -> None:
|
||||
events: Final = _events(("ocr", 0, None))
|
||||
gateway_results: Final = tuple(
|
||||
CaseResult(
|
||||
case=HarnessCase(
|
||||
strategy_id="trace_parity",
|
||||
strategy_label="Trace parity",
|
||||
sdk_function=sdk_function,
|
||||
spec=NotImplementedCaseSpec(reason=f"No {sdk_function} case is registered."),
|
||||
surface="gateway",
|
||||
),
|
||||
status=RunStatus.NOT_IMPLEMENTED,
|
||||
)
|
||||
for sdk_function in ("ocr", "messages")
|
||||
def test_renderer_groups_unavailable_entries_by_surface() -> None:
|
||||
gateway_result: Final = CaseResult(
|
||||
case=HarnessCase(
|
||||
strategy_id="trace_parity",
|
||||
strategy_label="Trace parity",
|
||||
sdk_function="messages",
|
||||
spec=NotImplementedCaseSpec(reason="No messages case is registered."),
|
||||
surface="gateway",
|
||||
),
|
||||
status=RunStatus.NOT_IMPLEMENTED,
|
||||
)
|
||||
|
||||
sections: Final = render_trace_results((_result(_comparison(events, events)), *gateway_results))
|
||||
sections: Final = render_trace_results((_result(_trace((), ())), gateway_result))
|
||||
|
||||
assert tuple(section.title for section in sections) == ("SDK trace comparisons", "GATEWAY trace comparisons")
|
||||
gateway_report: Final = "\n\n".join(sections[1].blocks)
|
||||
assert gateway_report.count("Not implemented") == 1
|
||||
assert "- ocr: No ocr case is registered." in gateway_report
|
||||
assert "- messages: No messages case is registered." in gateway_report
|
||||
assert tuple(section.title for section in sections) == ("SDK traces", "GATEWAY traces")
|
||||
assert "- messages: No messages case is registered." in "\n\n".join(sections[1].blocks)
|
||||
|
|
|
|||
|
|
@ -1,12 +1,23 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
import importlib
|
||||
import os
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Final, cast
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
from ...shared.reporting.models import Coverage, HarnessCase, HarnessRun, RunStatus, SdkFunction, Surface
|
||||
from ...shared.reporting.strategy import ModuleCaseSpec
|
||||
from ...shared.tracing.steps import Engine
|
||||
from ...shared.tracing.profiler import FunctionTraceEvent
|
||||
from ...shared.tracing.steps import Engine, PipelineStep, mapping
|
||||
from .models import GatewayRouteSpec, RouteFixture, RouteSpec, TraceScenario, TraceSuite
|
||||
from .runner import run_trace_mode, scenario_nodeids, validate_trace_suite
|
||||
from .reporting import TraceArtifact
|
||||
from .runner import run_trace_cases, run_trace_scenario, runner_selection, scenario_nodeids, validate_trace_suite
|
||||
from .sdk.execution import SdkCall, collect_trace, execute_trace
|
||||
|
||||
|
||||
def _fixture(_engine: Engine, _base_url: str) -> RouteFixture:
|
||||
|
|
@ -27,46 +38,243 @@ def test_scenario_filtering_and_occurrence_node_ids() -> None:
|
|||
suite: Final = TraceSuite(
|
||||
route=RouteSpec("ocr", ("ocr", "aocr"), ("ocr", "aocr"), _fixture),
|
||||
scenarios=(
|
||||
TraceScenario("one", _fixture, (), modes=("sync", "async")),
|
||||
TraceScenario("two", _fixture, (), modes=("async",)),
|
||||
TraceScenario("sync-one", _fixture, (), asynchronous=False),
|
||||
TraceScenario("async-one", _fixture, (), asynchronous=True),
|
||||
TraceScenario("async-two", _fixture, (), asynchronous=True),
|
||||
),
|
||||
)
|
||||
case: Final = _case()
|
||||
|
||||
nodes: Final = scenario_nodeids(suite, case, frozenset({"two"}))
|
||||
nodes: Final = scenario_nodeids(suite, case, frozenset({"async-two"}))
|
||||
|
||||
assert tuple(nodeid for _, _, nodeid in nodes) == ("trace:sdk:ocr:two:async",)
|
||||
assert tuple(nodeid for _, nodeid in nodes) == ("trace:sdk:ocr:async-two",)
|
||||
|
||||
|
||||
def test_python_engine_is_separate_from_scenario_selection() -> None:
|
||||
assert runner_selection(("mistral", "--engine=python")) == (frozenset({"mistral"}), "python")
|
||||
|
||||
|
||||
def test_python_engine_skips_native_bridge(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
|
||||
runner: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.runner")
|
||||
case: Final = _case()
|
||||
selected: list[tuple[frozenset[str], str]] = []
|
||||
|
||||
def reject_bridge(_repo_root: Path) -> str | None:
|
||||
raise AssertionError("Python-only tracing must not inspect or build the native bridge")
|
||||
|
||||
def capture_case(
|
||||
_run: HarnessRun,
|
||||
_case: HarnessCase,
|
||||
scenarios: frozenset[str],
|
||||
_on_update: object,
|
||||
engine: str,
|
||||
) -> None:
|
||||
selected.append((scenarios, engine))
|
||||
|
||||
monkeypatch.setattr(runner, "ensure_trace_bridge", reject_bridge)
|
||||
monkeypatch.setattr(runner, "_run_case", capture_case)
|
||||
|
||||
exit_code, _ = run_trace_cases((case,), tmp_path, lambda _: None, ("mistral", "--engine=python"))
|
||||
|
||||
assert exit_code == 0
|
||||
assert selected == [(frozenset({"mistral"}), "python")]
|
||||
|
||||
|
||||
def test_python_trace_preserves_native_ocr_dispatch_setting(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
execution: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.sdk.execution")
|
||||
route: Final = RouteSpec("ocr", ("ocr", "aocr"), ("ocr", "aocr"), _fixture)
|
||||
observed: list[str | None] = []
|
||||
|
||||
def collect(
|
||||
_function: SdkCall,
|
||||
_fixture: RouteFixture,
|
||||
_engine: Engine,
|
||||
*,
|
||||
asynchronous: bool,
|
||||
) -> SimpleNamespace:
|
||||
observed.append(os.environ.get("LITELLM_RUST"))
|
||||
return SimpleNamespace(
|
||||
events=(FunctionTraceEvent(0, None, "aocr" if asynchronous else "ocr"),),
|
||||
error=None,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(execution, "_collect", collect)
|
||||
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
collect_trace(route, "python", asynchronous=False)
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
collect_trace(route, "python", asynchronous=True)
|
||||
|
||||
assert observed == ["0", "1"]
|
||||
assert os.environ["LITELLM_RUST"] == "1"
|
||||
|
||||
|
||||
def test_expected_provider_failure_omits_feedback_banner(
|
||||
capsys: pytest.CaptureFixture[str], monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
loaded: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.sdk.responses.case")
|
||||
suite: Final = cast(TraceSuite, loaded.TRACE_SUITE)
|
||||
scenario: Final = next(item for item in suite.scenarios if item.name == "async-openai-provider-error")
|
||||
monkeypatch.setattr(litellm, "suppress_debug_info", False)
|
||||
assert isinstance(suite.route, RouteSpec)
|
||||
|
||||
result: Final = execute_trace(suite.route, scenario, "sdk", engine="python")
|
||||
|
||||
assert result.python_error is None
|
||||
assert "Give Feedback / Get Help" not in capsys.readouterr().out
|
||||
assert litellm.suppress_debug_info is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
def test_vertex_trace_keeps_unmapped_helpers_and_parents(asynchronous: bool) -> None:
|
||||
loaded: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.sdk.ocr.case")
|
||||
suite: Final = cast(TraceSuite, loaded.TRACE_SUITE)
|
||||
name: Final = f"{'async' if asynchronous else 'sync'}-vertex-deepseek"
|
||||
scenario: Final = next(item for item in suite.scenarios if item.name == name)
|
||||
assert isinstance(suite.route, RouteSpec)
|
||||
|
||||
trace: Final = execute_trace(suite.route, scenario, "sdk", engine="python")
|
||||
|
||||
assert trace.python_error is None
|
||||
url: Final = next(
|
||||
event for event in trace.python if event.raw.endswith(" VertexAIDeepSeekOCRConfig.get_complete_url")
|
||||
)
|
||||
project: Final = next(
|
||||
event for event in trace.python if event.raw.endswith(" VertexBase.safe_get_vertex_ai_project")
|
||||
)
|
||||
location: Final = next(
|
||||
event for event in trace.python if event.raw.endswith(" VertexBase.safe_get_vertex_ai_location")
|
||||
)
|
||||
assert project.parent_id == location.parent_id == url.id
|
||||
assert not any(event.raw.endswith(" VertexBase.get_access_token") for event in trace.python)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
def test_vertex_credentials_trace_runs_real_auth_helpers(asynchronous: bool, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
loaded: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.sdk.ocr.case")
|
||||
suite: Final = cast(TraceSuite, loaded.TRACE_SUITE)
|
||||
name: Final = f"{'async' if asynchronous else 'sync'}-vertex-deepseek-credentials"
|
||||
scenario: Final = next(item for item in suite.scenarios if item.name == name)
|
||||
monkeypatch.setenv("VERTEXAI_CREDENTIALS", "original-credentials")
|
||||
monkeypatch.setenv("VERTEX_AI_API_KEY", "original-api-key")
|
||||
assert isinstance(suite.route, RouteSpec)
|
||||
|
||||
trace: Final = execute_trace(suite.route, scenario, "sdk", engine="python")
|
||||
|
||||
assert trace.python_error is None
|
||||
validate: Final = next(
|
||||
event for event in trace.python if event.raw.endswith(" VertexAIDeepSeekOCRConfig.validate_environment")
|
||||
)
|
||||
helpers: Final = (
|
||||
"VertexBase.safe_get_vertex_ai_project",
|
||||
"VertexBase.safe_get_vertex_ai_credentials",
|
||||
"VertexBase.get_access_token",
|
||||
)
|
||||
assert tuple(event.raw.split(" ", 1)[1] for event in trace.python if event.parent_id == validate.id) == helpers
|
||||
token: Final = next(event for event in trace.python if event.raw.endswith(" VertexBase.get_access_token"))
|
||||
load: Final = next(event for event in trace.python if event.raw.endswith(" VertexBase.load_auth"))
|
||||
refresh: Final = next(event for event in trace.python if event.raw.endswith(" VertexBase.refresh_auth"))
|
||||
assert load.parent_id == token.id
|
||||
assert refresh.parent_id == load.id
|
||||
assert os.environ["VERTEXAI_CREDENTIALS"] == "original-credentials"
|
||||
assert os.environ["VERTEX_AI_API_KEY"] == "original-api-key"
|
||||
|
||||
|
||||
def test_gateway_trace_keeps_calls_outside_scenario_mappings(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
execution: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.gateway.execution")
|
||||
events: Final = (
|
||||
FunctionTraceEvent(0, None, "route.py:1 entry"),
|
||||
FunctionTraceEvent(1, 0, "auth.py:2 authenticate"),
|
||||
FunctionTraceEvent(2, 1, "auth.py:3 credentials"),
|
||||
)
|
||||
scenario: Final = TraceScenario(
|
||||
"async-gateway",
|
||||
_fixture,
|
||||
(mapping(rust_span="entry", python_frame=r" entry$"),),
|
||||
asynchronous=True,
|
||||
)
|
||||
monkeypatch.setattr(execution, "_collect", lambda *_args: events)
|
||||
|
||||
trace: Final = execution.execute_gateway_trace(GatewayRouteSpec("messages"), scenario, engine="python")
|
||||
|
||||
assert trace.python_error is None
|
||||
assert tuple((event.id, event.parent_id, event.raw) for event in trace.python) == tuple(
|
||||
(event.id, event.parent_id, event.raw) for event in events
|
||||
)
|
||||
|
||||
|
||||
def test_default_trace_skips_unavailable_rust_sdk_entrypoint(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
execution: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.sdk.execution")
|
||||
route: Final = RouteSpec("responses", ("responses", "aresponses"), None, _fixture)
|
||||
scenario: Final = TraceScenario("sync-openai", _fixture, (), asynchronous=False)
|
||||
engines: list[Engine] = []
|
||||
|
||||
def collect(_route: RouteSpec, engine: Engine, *, asynchronous: bool) -> tuple[FunctionTraceEvent, ...]:
|
||||
engines.append(engine)
|
||||
return (FunctionTraceEvent(0, None, "responses"),)
|
||||
|
||||
monkeypatch.setattr(execution, "collect_trace", collect)
|
||||
|
||||
trace: Final = execution.execute_trace(route, scenario, "sdk")
|
||||
|
||||
assert engines == ["python"]
|
||||
assert trace.engine == "python"
|
||||
assert trace.rust_error is None
|
||||
|
||||
|
||||
def test_default_trace_skips_unavailable_rust_gateway_route(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
execution: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.gateway.execution")
|
||||
route: Final = GatewayRouteSpec("responses", rust_supported=False)
|
||||
scenario: Final = TraceScenario("async-openai", _fixture, (), asynchronous=True)
|
||||
engines: list[Engine] = []
|
||||
|
||||
def collect(_route: GatewayRouteSpec, _scenario: TraceScenario, engine: Engine) -> tuple[FunctionTraceEvent, ...]:
|
||||
engines.append(engine)
|
||||
return (FunctionTraceEvent(0, None, "responses"),)
|
||||
|
||||
monkeypatch.setattr(execution, "_collect", collect)
|
||||
|
||||
trace: Final = execution.execute_gateway_trace(route, scenario)
|
||||
|
||||
assert engines == ["python"]
|
||||
assert trace.engine == "python"
|
||||
assert trace.rust_error is None
|
||||
|
||||
|
||||
def test_scenario_validation_rejects_duplicate_and_unsafe_names() -> None:
|
||||
route: Final = RouteSpec("ocr", ("ocr", "aocr"), ("ocr", "aocr"), _fixture)
|
||||
duplicate: Final = TraceSuite(
|
||||
route=route,
|
||||
scenarios=(TraceScenario("same", _fixture, ()), TraceScenario("same", _fixture, ())),
|
||||
scenarios=(
|
||||
TraceScenario("sync-same", _fixture, (), asynchronous=False),
|
||||
TraceScenario("sync-same", _fixture, (), asynchronous=False),
|
||||
),
|
||||
)
|
||||
unsafe: Final = TraceSuite(
|
||||
route=route, scenarios=(TraceScenario("sync-bad:name", _fixture, (), asynchronous=False),)
|
||||
)
|
||||
unsafe: Final = TraceSuite(route=route, scenarios=(TraceScenario("bad:name", _fixture, ()),))
|
||||
case: Final = _case()
|
||||
|
||||
assert validate_trace_suite(duplicate, case) is not None
|
||||
assert validate_trace_suite(unsafe, case) is not None
|
||||
|
||||
|
||||
def test_scenario_validation_rejects_invalid_modes_and_route_registration() -> None:
|
||||
invalid_modes: Final = TraceSuite(
|
||||
def test_scenario_validation_rejects_invalid_names_and_route_registration() -> None:
|
||||
invalid_name: Final = TraceSuite(
|
||||
route=RouteSpec("ocr", ("ocr", "aocr"), ("ocr", "aocr"), _fixture),
|
||||
scenarios=(TraceScenario("invalid", _fixture, (), modes=("sync", "sync")),),
|
||||
scenarios=(TraceScenario("bedrock", _fixture, (), asynchronous=True),),
|
||||
)
|
||||
wrong_function: Final = TraceSuite(
|
||||
route=RouteSpec("messages", ("create", "acreate"), ("messages", "amessages"), _fixture),
|
||||
scenarios=(TraceScenario("one", _fixture, ()),),
|
||||
scenarios=(TraceScenario("sync-one", _fixture, (), asynchronous=False),),
|
||||
)
|
||||
wrong_surface: Final = TraceSuite(
|
||||
route=GatewayRouteSpec("ocr"),
|
||||
scenarios=(TraceScenario("one", _fixture, ()),),
|
||||
scenarios=(TraceScenario("sync-one", _fixture, (), asynchronous=False),),
|
||||
)
|
||||
case: Final = _case()
|
||||
|
||||
assert "unique sync/async modes" in (validate_trace_suite(invalid_modes, case) or "")
|
||||
assert "start with sync- or async-" in (validate_trace_suite(invalid_name, case) or "")
|
||||
assert "does not match case function" in (validate_trace_suite(wrong_function, case) or "")
|
||||
assert "must use RouteSpec" in (validate_trace_suite(wrong_surface, case) or "")
|
||||
|
||||
|
|
@ -77,11 +285,35 @@ def test_invalid_route_dispatch_records_harness_error() -> None:
|
|||
result: Final = run.results[case.key]
|
||||
suite: Final = TraceSuite(
|
||||
route=GatewayRouteSpec("ocr"),
|
||||
scenarios=(TraceScenario("one", _fixture, (), modes=("sync",)),),
|
||||
scenarios=(TraceScenario("sync-one", _fixture, (), asynchronous=False),),
|
||||
)
|
||||
nodeid: Final = "trace:sdk:ocr:one:sync"
|
||||
nodeid: Final = "trace:sdk:ocr:sync-one"
|
||||
|
||||
run_trace_mode(run, result, suite, suite.scenarios[0], "sync", "sdk", nodeid, lambda _: None)
|
||||
run_trace_scenario(run, result, suite, suite.scenarios[0], "sdk", nodeid, lambda _: None)
|
||||
|
||||
assert result.outcomes[nodeid] is RunStatus.ERROR
|
||||
assert run.failures == [(nodeid, "gateway route cannot run on the sdk surface")]
|
||||
|
||||
|
||||
def test_different_python_and_rust_traces_pass(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
runner: Final = importlib.import_module("tests.rust-python-harness.strategies.trace_parity.runner")
|
||||
case: Final = _case()
|
||||
run: Final = HarnessRun.from_cases((case,))
|
||||
result: Final = run.results[case.key]
|
||||
suite: Final = TraceSuite(
|
||||
route=RouteSpec("ocr", ("ocr", "aocr"), ("ocr", "aocr"), _fixture),
|
||||
scenarios=(TraceScenario("sync-one", _fixture, (), asynchronous=False),),
|
||||
)
|
||||
trace: Final = TraceArtifact.from_traces(
|
||||
surface="sdk",
|
||||
sdk_function="ocr",
|
||||
scenario="sync-one",
|
||||
python=(PipelineStep(0, None, "python_step", "python.py:1 python_step"),),
|
||||
rust=(PipelineStep(0, None, "rust_step", "rust_step"),),
|
||||
)
|
||||
monkeypatch.setattr(runner, "_execute_scenario", lambda *_args: trace)
|
||||
|
||||
run_trace_scenario(run, result, suite, suite.scenarios[0], "sdk", "trace:sdk:ocr:sync-one", lambda _: None)
|
||||
|
||||
assert result.outcomes["trace:sdk:ocr:sync-one"] is RunStatus.PASSED
|
||||
assert run.failures == []
|
||||
|
|
|
|||
|
|
@ -38,30 +38,28 @@ def _trace_functions(
|
|||
python_functions: Final[dict[str, PythonFunctionIdentity]] = {}
|
||||
rust_functions: Final[dict[str, RustFunctionIdentity]] = {}
|
||||
for scenario in suite.scenarios:
|
||||
for mode in scenario.modes:
|
||||
route: Final = RouteSpec(
|
||||
route=suite.route.route,
|
||||
python_entrypoints=suite.route.python_entrypoints,
|
||||
rust_entrypoints=suite.route.rust_entrypoints,
|
||||
fixture=scenario.fixture,
|
||||
)
|
||||
python_trace: Final = collect_trace(route, "python", asynchronous=mode == "async")
|
||||
rust_trace: Final = collect_trace(route, "rust", asynchronous=mode == "async")
|
||||
if isinstance(python_trace, TraceExecutionFailure):
|
||||
raise ValueError(f"Python trace discovery failed for {scenario.name}/{mode}: {python_trace.message}")
|
||||
if isinstance(rust_trace, TraceExecutionFailure):
|
||||
raise ValueError(f"Rust trace discovery failed for {scenario.name}/{mode}: {rust_trace.message}")
|
||||
mappings: Final = scenario.mappings_for(mode)
|
||||
python_projection: Final = pipeline_projection("python", python_trace, mappings)
|
||||
rust_projection: Final = pipeline_projection("rust", rust_trace, mappings)
|
||||
for step in python_projection.steps:
|
||||
if step.span in spec.trace_spans:
|
||||
function: Final = PythonFunctionIdentity.from_trace(step.raw)
|
||||
python_functions[function.raw] = function
|
||||
for step in rust_projection.steps:
|
||||
if step.span in spec.trace_spans:
|
||||
function: Final = RustFunctionIdentity.from_trace(step.raw)
|
||||
rust_functions[step.raw] = function
|
||||
route: Final = RouteSpec(
|
||||
route=suite.route.route,
|
||||
python_entrypoints=suite.route.python_entrypoints,
|
||||
rust_entrypoints=suite.route.rust_entrypoints,
|
||||
fixture=scenario.fixture,
|
||||
)
|
||||
python_trace: Final = collect_trace(route, "python", asynchronous=scenario.asynchronous)
|
||||
rust_trace: Final = collect_trace(route, "rust", asynchronous=scenario.asynchronous)
|
||||
if isinstance(python_trace, TraceExecutionFailure):
|
||||
raise ValueError(f"Python trace discovery failed for {scenario.name}: {python_trace.message}")
|
||||
if isinstance(rust_trace, TraceExecutionFailure):
|
||||
raise ValueError(f"Rust trace discovery failed for {scenario.name}: {rust_trace.message}")
|
||||
python_projection: Final = pipeline_projection("python", python_trace, scenario.mappings)
|
||||
rust_projection: Final = pipeline_projection("rust", rust_trace, scenario.mappings)
|
||||
for step in python_projection.steps:
|
||||
if step.span in spec.trace_spans:
|
||||
function: Final = PythonFunctionIdentity.from_trace(step.raw)
|
||||
python_functions[function.raw] = function
|
||||
for step in rust_projection.steps:
|
||||
if step.span in spec.trace_spans:
|
||||
function: Final = RustFunctionIdentity.from_trace(step.raw)
|
||||
rust_functions[step.raw] = function
|
||||
if not python_functions or not rust_functions:
|
||||
raise ValueError(f"Python trace discovery found no functions for spans: {', '.join(spec.trace_spans)}")
|
||||
return (
|
||||
|
|
|
|||
|
|
@ -1,9 +1,12 @@
|
|||
import logging
|
||||
import re
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.caching.redis_cache as redis_cache_module
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.caching.redis_cache import RedisCache, _RedisTimeoutLogThrottle
|
||||
from litellm.types.caching import LiteLLMCacheType, SemanticCacheScope
|
||||
from litellm.types.utils import Embedding, EmbeddingResponse, Usage
|
||||
|
||||
|
|
@ -53,6 +56,29 @@ def test_cache_key_debug_log_does_not_include_prompt_material(caplog):
|
|||
assert any(cache_key in message for message in created_cache_key_logs)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("backend", "expected_level"),
|
||||
[
|
||||
pytest.param(MagicMock(spec=RedisCache), logging.DEBUG, id="redis_backend_is_throttled"),
|
||||
pytest.param(MagicMock(), logging.ERROR, id="other_backend_logs_every_timeout"),
|
||||
],
|
||||
)
|
||||
def test_add_cache_timeout_only_joins_redis_throttle_for_redis_backends(backend, expected_level, caplog, monkeypatch):
|
||||
throttle = _RedisTimeoutLogThrottle(interval=5.0, clock=MagicMock(return_value=1_000.0))
|
||||
assert throttle.admit() == 0
|
||||
monkeypatch.setattr(redis_cache_module, "_redis_timeout_log_throttle", throttle)
|
||||
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
backend.set_cache.side_effect = TimeoutError("lit7520 backend timed out")
|
||||
cache.cache = backend
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
cache.add_cache("result", model="gpt-4.1-mini", messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
records = [r for r in caplog.records if "lit7520 backend timed out" in r.getMessage()]
|
||||
assert [r.levelno for r in records] == [expected_level]
|
||||
|
||||
|
||||
def _embedding_response(prompt_tokens, num_items):
|
||||
return EmbeddingResponse(
|
||||
model="amazon.titan-embed-image-v1",
|
||||
|
|
|
|||
|
|
@ -704,3 +704,58 @@ async def test_open_breaker_keeps_async_batch_read_memory_hits_and_releases_rese
|
|||
|
||||
assert list(await cache.async_batch_get_cache(["k1", "k2"])) == ["v1", None]
|
||||
assert "k2" not in cache.last_redis_batch_access_time
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redis_timeouts_falling_back_to_memory_log_once_per_interval(caplog, monkeypatch):
|
||||
"""The first fallback WARNING of a timeout streak logs, the rest stay at DEBUG until the summary."""
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
from litellm.caching import redis_cache as redis_cache_module
|
||||
from litellm.caching.redis_cache import _RedisTimeoutLogThrottle
|
||||
|
||||
clock = MagicMock(return_value=1_000.0)
|
||||
monkeypatch.setattr(
|
||||
redis_cache_module, "_redis_timeout_log_throttle", _RedisTimeoutLogThrottle(interval=5.0, clock=clock)
|
||||
)
|
||||
|
||||
class _TimingOutRedis:
|
||||
async def async_increment_pipeline(self, increment_list, **kwargs):
|
||||
raise RedisTimeoutError("Timeout reading from 127.0.0.1:6379")
|
||||
|
||||
async def async_increment(self, key, value, **kwargs):
|
||||
raise RedisTimeoutError("Timeout reading from 127.0.0.1:6379")
|
||||
|
||||
cache = DualCache(
|
||||
in_memory_cache=InMemoryCache(),
|
||||
redis_cache=_TimingOutRedis(), # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
)
|
||||
increments = [RedisPipelineIncrementOperation(key="k", increment_value=1.0, ttl=60)]
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
for _ in range(100):
|
||||
await cache.async_increment_cache_pipeline(increment_list=increments)
|
||||
await cache.async_increment_cache("k", 1.0)
|
||||
|
||||
visible = [r for r in caplog.records if r.levelno >= logging.WARNING]
|
||||
assert [(r.levelno, r.getMessage()) for r in visible] == [
|
||||
(
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache_pipeline failed, falling back to in-memory result:"
|
||||
" Timeout reading from 127.0.0.1:6379",
|
||||
)
|
||||
]
|
||||
assert visible[0].filename == "dual_cache.py"
|
||||
assert sum("Timeout reading from" in r.getMessage() for r in caplog.records) == 200
|
||||
|
||||
caplog.clear()
|
||||
clock.return_value += 5.0
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
await cache.async_increment_cache("k", 1.0)
|
||||
assert [(r.levelno, r.getMessage()) for r in caplog.records] == [
|
||||
(
|
||||
logging.WARNING,
|
||||
"Redis async_increment_cache failed, falling back to in-memory result: Timeout reading from 127.0.0.1:6379"
|
||||
" (199 more Redis timeouts since the previous Redis timeout line were logged at DEBUG)",
|
||||
)
|
||||
]
|
||||
|
|
|
|||
|
|
@ -978,17 +978,17 @@ async def test_stale_timeout_does_not_let_sub_threshold_hard_failures_open_the_b
|
|||
from redis.exceptions import ConnectionError as RedisConnectionError
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, _is_redis_timeout_failure
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, is_redis_timeout_failure
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=0.05)
|
||||
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
await asyncio.sleep(0.06)
|
||||
for _ in range(breaker.failure_threshold - 1):
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
assert breaker.is_open() is False, "2 hard failures and 1 stale timeout are below both thresholds"
|
||||
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
assert breaker.is_open() is True, "the threshold-th hard failure must still open it"
|
||||
|
||||
|
||||
|
|
@ -1000,19 +1000,19 @@ async def test_hard_failure_resets_timeout_streak_so_a_later_burst_must_earn_its
|
|||
from redis.exceptions import ConnectionError as RedisConnectionError
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, _is_redis_timeout_failure
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, is_redis_timeout_failure
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=0.05)
|
||||
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
await asyncio.sleep(0.06)
|
||||
for _ in range(breaker.failure_threshold):
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
assert breaker.is_open() is False, "the burst is instantaneous, so the duration gate must hold it closed"
|
||||
|
||||
await asyncio.sleep(0.06)
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisTimeoutError("read timed out")))
|
||||
assert breaker.is_open() is True, "the same run of timeouts persisting past the duration must open it"
|
||||
|
||||
|
||||
|
|
@ -1023,7 +1023,7 @@ async def test_breaker_metrics_track_state_and_failure_class():
|
|||
from redis.exceptions import ConnectionError as RedisConnectionError
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, _is_redis_timeout_failure
|
||||
from litellm.caching.redis_cache import RedisCircuitBreaker, is_redis_timeout_failure
|
||||
|
||||
def sample(name, labels=None):
|
||||
return REGISTRY.get_sample_value(name, labels) or 0.0
|
||||
|
|
@ -1035,9 +1035,9 @@ async def test_breaker_metrics_track_state_and_failure_class():
|
|||
closed_gauge_before = sample("litellm_redis_circuit_breaker_state", {"state": "closed"})
|
||||
|
||||
breaker = RedisCircuitBreaker(failure_threshold=2, recovery_timeout=60, timeout_min_duration=5.0)
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("t")))
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisTimeoutError("t")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
breaker.record_failure(is_timeout=is_redis_timeout_failure(RedisConnectionError("refused")))
|
||||
|
||||
assert sample("litellm_redis_circuit_breaker_failures_total", {"failure_class": "timeout"}) == timeout_before + 1
|
||||
assert sample("litellm_redis_circuit_breaker_failures_total", {"failure_class": "connectivity"}) == hard_before + 2
|
||||
|
|
@ -1205,6 +1205,117 @@ async def test_a_probe_overtaken_by_a_later_outage_leaves_the_breaker_to_the_new
|
|||
assert breaker._state == breaker.CLOSED
|
||||
|
||||
|
||||
def test_timeouts_during_a_blip_log_once_per_interval_not_once_per_call(sync_batch_redis_cache, caplog, monkeypatch):
|
||||
"""A timeout streak logs its first failure plus one summary per interval; other failures log per call."""
|
||||
import logging
|
||||
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
from litellm.caching import redis_cache as redis_cache_module
|
||||
from litellm.caching.redis_cache import _RedisTimeoutLogThrottle
|
||||
|
||||
clock = MagicMock(return_value=1_000.0)
|
||||
monkeypatch.setattr(
|
||||
redis_cache_module, "_redis_timeout_log_throttle", _RedisTimeoutLogThrottle(interval=5.0, clock=clock)
|
||||
)
|
||||
sync_batch_redis_cache.redis_client.get.side_effect = RedisTimeoutError("Timeout reading from 127.0.0.1:6379")
|
||||
sync_batch_redis_cache.redis_client.mget.side_effect = RedisTimeoutError("Timeout reading from 127.0.0.1:6379")
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
for _ in range(200):
|
||||
assert sync_batch_redis_cache.get_cache("lit7520") is None
|
||||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit7520"]) == {}
|
||||
|
||||
timeout_records = [r for r in caplog.records if "Timeout reading from" in r.getMessage()]
|
||||
assert len(timeout_records) == 201, "every timeout must stay visible at DEBUG"
|
||||
assert [r.getMessage() for r in timeout_records if r.levelno >= logging.WARNING] == [
|
||||
"litellm.caching.caching: get() - Got exception from REDIS: Timeout reading from 127.0.0.1:6379"
|
||||
]
|
||||
assert timeout_records[0].levelno == logging.ERROR
|
||||
assert timeout_records[0].filename == "redis_cache.py"
|
||||
assert timeout_records[0].lineno != timeout_records[-1].lineno, "the record must point at the cache operation"
|
||||
|
||||
caplog.clear()
|
||||
clock.return_value += 5.0
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit7520"]) == {}
|
||||
assert [(r.levelno, r.getMessage()) for r in caplog.records] == [
|
||||
(
|
||||
logging.ERROR,
|
||||
"Error occurred in batch get cache: Timeout reading from 127.0.0.1:6379"
|
||||
" (200 more Redis timeouts since the previous Redis timeout line were logged at DEBUG)",
|
||||
)
|
||||
]
|
||||
|
||||
caplog.clear()
|
||||
sync_batch_redis_cache.redis_client.get.side_effect = OSError("redis unavailable")
|
||||
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
|
||||
for _ in range(3):
|
||||
assert sync_batch_redis_cache.get_cache("lit7520") is None
|
||||
assert [r.levelno for r in caplog.records if "redis unavailable" in r.getMessage()] == [logging.ERROR] * 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"call_method",
|
||||
[
|
||||
pytest.param(lambda c: c.async_set_cache_pipeline([("lit7520", "v")]), id="async_set_cache_pipeline"),
|
||||
pytest.param(
|
||||
lambda c: c.async_set_cache_pipeline_with_ttls([("lit7520", "v", 60.0)]),
|
||||
id="async_set_cache_pipeline_with_ttls",
|
||||
),
|
||||
pytest.param(lambda c: c.async_set_cache_sadd("lit7520", ["v"], ttl=None), id="async_set_cache_sadd"),
|
||||
pytest.param(lambda c: c.async_increment("lit7520", 1.0), id="async_increment"),
|
||||
pytest.param(
|
||||
lambda c: c.async_increment_pipeline([{"key": "lit7520", "increment_value": 1.0, "ttl": 60}]),
|
||||
id="async_increment_pipeline",
|
||||
),
|
||||
pytest.param(lambda c: c.async_rpush("lit7520", ["v"]), id="async_rpush"),
|
||||
pytest.param(
|
||||
lambda c: c.async_rpush_pipeline([{"key": "lit7520", "values": ["v"]}]), id="async_rpush_pipeline"
|
||||
),
|
||||
pytest.param(lambda c: c.async_lpop("lit7520"), id="async_lpop"),
|
||||
pytest.param(lambda c: c.async_lpop_pipeline([{"key": "lit7520", "count": 1}]), id="async_lpop_pipeline"),
|
||||
],
|
||||
)
|
||||
async def test_write_path_timeouts_inside_the_interval_stay_at_debug(call_method, caplog, monkeypatch, redis_no_ping):
|
||||
"""A write or list operation timing out mid-streak is counted by the throttle instead of logging its own ERROR."""
|
||||
import contextlib
|
||||
import logging
|
||||
|
||||
from redis.exceptions import TimeoutError as RedisTimeoutError
|
||||
|
||||
from litellm.caching import redis_cache as redis_cache_module
|
||||
from litellm.caching.redis_cache import _RedisTimeoutLogThrottle
|
||||
|
||||
clock = MagicMock(return_value=1_000.0)
|
||||
throttle = _RedisTimeoutLogThrottle(interval=5.0, clock=clock)
|
||||
assert throttle.admit() == 0
|
||||
monkeypatch.setattr(redis_cache_module, "_redis_timeout_log_throttle", throttle)
|
||||
|
||||
timeout = RedisTimeoutError("Timeout reading from 127.0.0.1:6379")
|
||||
client = MagicMock()
|
||||
client.pipeline.return_value.__aenter__.side_effect = timeout
|
||||
client.sadd = AsyncMock(side_effect=timeout)
|
||||
client.incrbyfloat = AsyncMock(side_effect=timeout)
|
||||
client.rpush = AsyncMock(side_effect=timeout)
|
||||
client.lpop = AsyncMock(side_effect=timeout)
|
||||
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
|
||||
cache = RedisCache()
|
||||
|
||||
with (
|
||||
patch.object(cache, "init_async_client", return_value=client),
|
||||
caplog.at_level(logging.DEBUG, logger="LiteLLM"),
|
||||
):
|
||||
with contextlib.suppress(RedisTimeoutError):
|
||||
await call_method(cache)
|
||||
|
||||
timeout_records = [r for r in caplog.records if "Timeout reading from" in r.getMessage()]
|
||||
assert [(r.levelno, r.filename) for r in timeout_records] == [(logging.DEBUG, "redis_cache.py")]
|
||||
clock.return_value += 5.0
|
||||
assert throttle.admit() == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pool_wait_timeout_is_a_timeout_failure_not_hard_connectivity():
|
||||
"""A saturated blocking pool must not open the breaker before the timeout minimum duration.
|
||||
|
|
@ -1241,7 +1352,7 @@ async def test_pool_wait_timeout_is_a_timeout_failure_not_hard_connectivity():
|
|||
def test_timeout_classification_follows_the_explicit_cause_chain_only():
|
||||
from redis.exceptions import ConnectionError as RedisConnectionError
|
||||
|
||||
from litellm.caching.redis_cache import _is_redis_timeout_failure
|
||||
from litellm.caching.redis_cache import is_redis_timeout_failure
|
||||
|
||||
def raise_chained_from_timeout() -> None:
|
||||
try:
|
||||
|
|
@ -1260,9 +1371,9 @@ def test_timeout_classification_follows_the_explicit_cause_chain_only():
|
|||
with pytest.raises(RedisConnectionError) as contextual:
|
||||
raise_while_handling_timeout()
|
||||
|
||||
assert _is_redis_timeout_failure(chained.value) is True
|
||||
assert _is_redis_timeout_failure(contextual.value) is False
|
||||
assert _is_redis_timeout_failure(RedisConnectionError("refused")) is False
|
||||
assert is_redis_timeout_failure(chained.value) is True
|
||||
assert is_redis_timeout_failure(contextual.value) is False
|
||||
assert is_redis_timeout_failure(RedisConnectionError("refused")) is False
|
||||
|
||||
|
||||
class _RoundTripCountingRedis:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
"""Golden tests for the OTel v2 engine: span shape, kinds, semconv attributes,
|
||||
legacy dual-emit, hierarchy, error status, and idempotency. Needs the OTel SDK."""
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("opentelemetry")
|
||||
|
|
@ -18,6 +20,7 @@ from litellm.integrations.otel.plumbing import providers # noqa: E402
|
|||
from litellm.integrations.otel.emitter import SpanEmitter # noqa: E402
|
||||
from litellm.integrations.otel.emitter import stamp_error # noqa: E402
|
||||
from litellm.integrations.otel.mappers.utils import ( # noqa: E402
|
||||
MAX_MESSAGE_ATTRS_PER_SPAN,
|
||||
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
|
||||
)
|
||||
from litellm.integrations.otel.model.payloads import ( # noqa: E402
|
||||
|
|
@ -440,3 +443,161 @@ def test_vendor_tool_definitions_are_truncated_not_dropped():
|
|||
assert a["llm.tools.0.tool.name"] == "tool_0"
|
||||
assert a["llm.tools.0.tool.json_schema"]
|
||||
assert "llm.tools.126.tool.name" not in a
|
||||
|
||||
|
||||
def _conversation_payload(turns, choices=1, **overrides):
|
||||
"""A ``turns``-message chat with ``choices`` response choices, content-bearing."""
|
||||
return _payload(
|
||||
messages=[{"role": ("user", "assistant")[i % 2], "content": f"turn {i}"} for i in range(turns)],
|
||||
response={
|
||||
"id": "resp_1",
|
||||
"model": "gpt-4o-2024",
|
||||
"choices": [
|
||||
{"finish_reason": "stop", "message": {"role": "assistant", "content": f"reply {i}"}}
|
||||
for i in range(choices)
|
||||
],
|
||||
},
|
||||
**overrides,
|
||||
)
|
||||
|
||||
|
||||
def _conversation_span(mapper_names, payload, legacy_compat=False):
|
||||
"""The exported LLM-call span for ``payload`` with content capture on."""
|
||||
cfg = OpenTelemetryV2Config(
|
||||
exporter="in_memory",
|
||||
legacy_compat=legacy_compat,
|
||||
mapper_names=list(mapper_names),
|
||||
capture_message_content="span_only",
|
||||
)
|
||||
provider, exporter = providers.in_memory_provider(cfg)
|
||||
engine = SpanEmitter(providers.get_tracer(provider, "litellm-test"), cfg)
|
||||
engine.emit(
|
||||
SpanRole.LLM_CALL,
|
||||
LLMCallSpanData.from_standard_logging_payload(payload, capture_content=True),
|
||||
)
|
||||
(span,) = exporter.get_finished_spans()
|
||||
return span
|
||||
|
||||
|
||||
def _indexed_message_count(attributes, prefix):
|
||||
return len({key.split(".")[2] for key in attributes if key.startswith(f"{prefix}.")})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("turns", [60, 200])
|
||||
def test_long_conversation_does_not_evict_core_attributes(turns):
|
||||
"""Per-message OpenInference attributes must never crowd core telemetry off the span."""
|
||||
span = _conversation_span(["genai", "openinference"], _conversation_payload(turns))
|
||||
a = span.attributes
|
||||
|
||||
assert span.dropped_attributes == 0
|
||||
assert a[GenAI.REQUEST_MODEL] == "gpt-4o"
|
||||
assert a[GenAI.PROVIDER_NAME] == "openai"
|
||||
assert a[GenAI.USAGE_INPUT_TOKENS] == 10
|
||||
assert a[GenAI.USAGE_OUTPUT_TOKENS] == 5
|
||||
assert a[GenAI.RESPONSE_FINISH_REASONS] == ("stop",)
|
||||
assert a[f"{LiteLLM.COST_PREFIX}total"] == 0.002
|
||||
|
||||
assert a["llm.input_messages.0.message.content"] == "turn 0"
|
||||
assert a["llm.output_messages.0.message.content"] == "reply 0"
|
||||
assert a[f"llm.input_messages.{turns - 1}.message.content"] == f"turn {turns - 1}"
|
||||
assert f"llm.input_messages.{turns // 2}.message.role" not in a
|
||||
assert len(json.loads(a["input.value"])) == turns
|
||||
assert len(json.loads(a["output.value"])) == 1
|
||||
assert len(json.loads(a[GenAI.INPUT_MESSAGES])) == turns
|
||||
|
||||
|
||||
def test_short_conversation_keeps_every_message_indexed():
|
||||
"""Below the cap nothing is truncated in either direction."""
|
||||
a = _conversation_span(["genai", "openinference"], _conversation_payload(4, choices=2)).attributes
|
||||
for idx in range(4):
|
||||
assert a[f"llm.input_messages.{idx}.message.content"] == f"turn {idx}"
|
||||
for idx in range(2):
|
||||
assert a[f"llm.output_messages.{idx}.message.content"] == f"reply {idx}"
|
||||
|
||||
|
||||
def test_indexed_prompt_keeps_opener_and_latest_turns_under_a_value_length_limit(monkeypatch):
|
||||
"""The system prompt and the live turn keep their own keys once the SDK clips ``input.value``."""
|
||||
monkeypatch.setenv("OTEL_SPAN_ATTRIBUTE_VALUE_LENGTH_LIMIT", "256")
|
||||
chat = _conversation_payload(60)
|
||||
payload = {
|
||||
**chat,
|
||||
"messages": [
|
||||
{"role": "system", "content": "be terse"},
|
||||
*chat["messages"][1:-1],
|
||||
{"role": "user", "content": "LATEST-TURN"},
|
||||
],
|
||||
}
|
||||
a = _conversation_span(["genai", "openinference"], payload).attributes
|
||||
|
||||
assert len(a["input.value"]) == 256
|
||||
assert a["llm.input_messages.0.message.role"] == "system"
|
||||
assert a["llm.input_messages.0.message.content"] == "be terse"
|
||||
assert a["llm.input_messages.59.message.role"] == "user"
|
||||
assert a["llm.input_messages.59.message.content"] == "LATEST-TURN"
|
||||
assert a["llm.output_messages.0.message.content"] == "reply 0"
|
||||
assert [int(key.split(".")[2]) for key in a if key.endswith("message.content") and key.startswith("llm.input_")] == [
|
||||
0,
|
||||
*range(54, 60),
|
||||
]
|
||||
|
||||
|
||||
def test_message_cap_is_shared_across_input_and_output():
|
||||
"""One span-wide allowance covers both directions, and the response always keeps a share."""
|
||||
long_prompt = _conversation_span(["genai", "openinference"], _conversation_payload(60, choices=1)).attributes
|
||||
many_choices = _conversation_span(["genai", "openinference"], _conversation_payload(60, choices=20)).attributes
|
||||
|
||||
single_reply_indexed = _indexed_message_count(long_prompt, "llm.output_messages")
|
||||
assert single_reply_indexed == 1
|
||||
assert _indexed_message_count(long_prompt, "llm.input_messages") + single_reply_indexed == (
|
||||
MAX_MESSAGE_ATTRS_PER_SPAN // 2
|
||||
)
|
||||
|
||||
assert _indexed_message_count(many_choices, "llm.input_messages") > 0
|
||||
assert _indexed_message_count(many_choices, "llm.output_messages") > single_reply_indexed
|
||||
assert _indexed_message_count(many_choices, "llm.input_messages") + _indexed_message_count(
|
||||
many_choices, "llm.output_messages"
|
||||
) == (MAX_MESSAGE_ATTRS_PER_SPAN // 2)
|
||||
|
||||
|
||||
def test_fully_populated_span_with_every_vocabulary_stays_within_the_attribute_limit():
|
||||
"""Every capped family maxed at once still leaves the whole core intact."""
|
||||
payload = _conversation_payload(
|
||||
200,
|
||||
choices=20,
|
||||
stream=True,
|
||||
model_parameters={
|
||||
**_tools_payload(127)["model_parameters"],
|
||||
"top_p": 0.9,
|
||||
"frequency_penalty": 0.1,
|
||||
"presence_penalty": 0.1,
|
||||
"seed": 7,
|
||||
"stop": ["\n"],
|
||||
},
|
||||
cost_breakdown={
|
||||
key: 0.001
|
||||
for key in (
|
||||
"input_cost",
|
||||
"output_cost",
|
||||
"cache_read_cost",
|
||||
"cache_creation_cost",
|
||||
"tool_usage_cost",
|
||||
"original_cost",
|
||||
"discount_amount",
|
||||
"discount_percent",
|
||||
"margin_fixed_amount",
|
||||
"margin_percent",
|
||||
"margin_total_amount",
|
||||
"total_cost",
|
||||
)
|
||||
},
|
||||
)
|
||||
span = _conversation_span(["genai", "openinference", "langfuse", "weave", "langtrace"], payload, legacy_compat=True)
|
||||
a = span.attributes
|
||||
|
||||
assert span.dropped_attributes == 0
|
||||
assert a[GenAI.REQUEST_MODEL] == "gpt-4o"
|
||||
assert a[f"{LiteLLM.COST_PREFIX}total"] == 0.002
|
||||
assert a[LiteLLM.TOOLS_DECLARED] == 127
|
||||
assert a["llm.input_messages.0.message.content"] == "turn 0"
|
||||
assert a["llm.input_messages.199.message.content"] == "turn 199"
|
||||
assert a["llm.output_messages.0.message.content"] == "reply 0"
|
||||
|
|
|
|||
|
|
@ -716,6 +716,77 @@ async def test_failure_hook_emits_api_provider_value_on_failed_requests_metric()
|
|||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
async def _failed_requests_api_provider_labels(
|
||||
request_data: dict[str, object],
|
||||
original_exception: Exception,
|
||||
) -> list[str]:
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
await PrometheusLogger().async_post_call_failure_hook(
|
||||
request_data=request_data,
|
||||
original_exception=original_exception,
|
||||
user_api_key_dict=UserAPIKeyAuth(token="tok"),
|
||||
)
|
||||
return [
|
||||
s.labels.get("api_provider")
|
||||
for s in _collected_samples("litellm_proxy_failed_requests_metric_total")
|
||||
]
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_hook_emits_api_provider_from_pre_call_rate_limit_error_for_router_alias():
|
||||
"""
|
||||
Pre-call limiters reject before a deployment lands on request_data and a
|
||||
router alias cannot be inferred from its name, so the provider the limiter
|
||||
resolved onto the exception is the only source for the label.
|
||||
"""
|
||||
from litellm.exceptions import RateLimitType
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
err = ProxyRateLimitError(
|
||||
detail={"error": "rpm exceeded"},
|
||||
rate_limit_type=RateLimitType.REQUESTS,
|
||||
model="openai/gpt-5.4-mini",
|
||||
llm_provider="openai",
|
||||
)
|
||||
|
||||
assert await _failed_requests_api_provider_labels(
|
||||
{"model": "team-chat-model", "metadata": {}}, err
|
||||
) == ["openai"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_hook_leaves_api_provider_unset_when_rate_limiter_could_not_resolve_provider():
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
err = ProxyRateLimitError(detail={"error": "rpm exceeded"}, model="unknown-alias")
|
||||
|
||||
assert await _failed_requests_api_provider_labels(
|
||||
{"model": "unknown-alias", "metadata": {}}, err
|
||||
) == ["None"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_hook_prefers_request_data_provider_over_exception_provider():
|
||||
from litellm.exceptions import RateLimitError
|
||||
|
||||
err = RateLimitError(message="upstream 429", llm_provider="openai", model="gpt-4o")
|
||||
|
||||
assert await _failed_requests_api_provider_labels(
|
||||
{
|
||||
"model": "gpt-4o",
|
||||
"metadata": {},
|
||||
"litellm_params": {"custom_llm_provider": "azure"},
|
||||
},
|
||||
err,
|
||||
) == ["azure"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_user_email_in_required_metrics()
|
||||
test_user_email_label_exists()
|
||||
|
|
|
|||
|
|
@ -2101,6 +2101,37 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
|
|||
assert completion_cost == pytest.approx(zone_multiplier * output_multiplier * completion_tokens * 5e-5)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,input_rate,cache_read_rate,output_rate",
|
||||
[
|
||||
("azure/gpt-chat-latest", 5e-6, 5e-7, 3e-5),
|
||||
("azure/chat-latest", 5e-6, 5e-7, 3e-5),
|
||||
("azure/us/gpt-chat-latest", 5.5e-6, 5.5e-7, 3.3e-5),
|
||||
],
|
||||
)
|
||||
def test_generic_cost_per_token_azure_gpt_chat_latest_price_sheet(
|
||||
_local_model_cost_map, model, input_rate, cache_read_rate, output_rate
|
||||
):
|
||||
"""The Azure OpenAI price sheet lists GPT-Chat Latest at $5 input, $0.50 cached input and $30 output per 1M
|
||||
tokens on Global, and $5.50, $0.55 and $33 on Data Zone. Foundry names the product gpt-chat-latest and the
|
||||
OpenAI API names the same model chat-latest, so both spellings bill the Global sheet.
|
||||
"""
|
||||
prompt_tokens = 100000
|
||||
cached_tokens = 40000
|
||||
completion_tokens = 1000
|
||||
usage = Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="azure")
|
||||
|
||||
assert prompt_cost == pytest.approx((prompt_tokens - cached_tokens) * input_rate + cached_tokens * cache_read_rate)
|
||||
assert completion_cost == pytest.approx(completion_tokens * output_rate)
|
||||
|
||||
|
||||
def test_generic_cost_per_token_azure_ai_gpt_6_astra_flex_bills_the_standard_rate(_local_model_cost_map):
|
||||
usage = Usage(prompt_tokens=1000, completion_tokens=100, total_tokens=1100)
|
||||
|
||||
|
|
@ -2648,6 +2679,7 @@ def test_cache_writing_cost_with_zero_creation_tokens_and_ephemeral_details():
|
|||
|
||||
prompt_tokens_details: PromptTokensDetailsResult = {
|
||||
"cache_hit_tokens": 0,
|
||||
"cache_hit_audio_tokens": 0,
|
||||
"cache_creation_tokens": 0,
|
||||
"cache_creation_token_details": CacheCreationTokenDetails(
|
||||
ephemeral_5m_input_tokens=100,
|
||||
|
|
@ -4005,6 +4037,7 @@ def test_billed_token_rates_follow_the_token_tier_the_breakdown_bills_at(monkeyp
|
|||
input_cost_per_token=6e-6,
|
||||
output_cost_per_token=3e-5,
|
||||
cache_read_input_token_cost=6e-7,
|
||||
cache_read_input_audio_token_cost=6e-7,
|
||||
cache_creation_input_token_cost=7.5e-6,
|
||||
cache_creation_input_token_cost_above_1hr=0.0,
|
||||
output_cost_per_reasoning_token=3e-5,
|
||||
|
|
@ -5147,3 +5180,157 @@ def test_generic_cost_per_token_bills_nested_reasoning_once_beside_audio_output(
|
|||
assert completion_cost == pytest.approx(
|
||||
30 * info["output_cost_per_token"] + 70 * info["output_cost_per_audio_token"]
|
||||
)
|
||||
|
||||
|
||||
def test_cached_realtime_audio_tokens_billed_at_audio_cache_read_rate(
|
||||
_local_model_cost_map: None,
|
||||
) -> None:
|
||||
usage = Usage(
|
||||
prompt_tokens=283,
|
||||
completion_tokens=0,
|
||||
total_tokens=283,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=116,
|
||||
audio_tokens=167,
|
||||
cached_tokens=192,
|
||||
cached_tokens_details={"text_tokens": 64, "audio_tokens": 128},
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gpt-realtime-2", usage=usage, custom_llm_provider="openai"
|
||||
)
|
||||
assert prompt_cost == pytest.approx(0.0015328)
|
||||
|
||||
|
||||
def test_prompt_tokens_details_without_cached_tokens_details_unchanged(
|
||||
_local_model_cost_map: None,
|
||||
) -> None:
|
||||
usage = Usage(
|
||||
prompt_tokens=283,
|
||||
completion_tokens=0,
|
||||
total_tokens=283,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=116, audio_tokens=167, cached_tokens=192
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gpt-realtime-2", usage=usage, custom_llm_provider="openai"
|
||||
)
|
||||
assert prompt_cost == pytest.approx(0.0029888)
|
||||
|
||||
|
||||
def test_cached_audio_tokens_fall_back_to_cache_read_input_token_cost() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"input_cost_per_token": 4e-6,
|
||||
"input_cost_per_audio_token": 32e-6,
|
||||
"cache_read_input_token_cost": 5e-7,
|
||||
}
|
||||
usage = Usage(
|
||||
prompt_tokens=283,
|
||||
completion_tokens=0,
|
||||
total_tokens=283,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=116,
|
||||
audio_tokens=167,
|
||||
cached_tokens=192,
|
||||
cached_tokens_details={"text_tokens": 64, "audio_tokens": 128},
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="some-realtime-model",
|
||||
usage=usage,
|
||||
custom_llm_provider="openai",
|
||||
model_info=model_info,
|
||||
)
|
||||
expected = 52 * 4e-6 + 64 * 5e-7 + 39 * 32e-6 + 128 * 5e-7
|
||||
assert prompt_cost == pytest.approx(expected)
|
||||
|
||||
|
||||
def test_cached_audio_tokens_capped_at_cached_tokens(_local_model_cost_map: None) -> None:
|
||||
"""Nested cached_tokens_details exceeding cached_tokens must not over-subtract the audio bucket."""
|
||||
usage = Usage(
|
||||
prompt_tokens=283,
|
||||
completion_tokens=0,
|
||||
total_tokens=283,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=116,
|
||||
audio_tokens=167,
|
||||
cached_tokens=100,
|
||||
cached_tokens_details={"audio_tokens": 128},
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gpt-realtime-2", usage=usage, custom_llm_provider="openai"
|
||||
)
|
||||
assert prompt_cost == pytest.approx(116 * 4e-6 + (167 - 100) * 32e-6 + 100 * 4e-7)
|
||||
|
||||
|
||||
def test_cached_audio_tokens_billed_at_audio_cache_rate_through_model_info_lookup(_local_model_cost_map: None) -> None:
|
||||
usage = Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=0,
|
||||
total_tokens=1000,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=400,
|
||||
audio_tokens=600,
|
||||
cached_tokens=500,
|
||||
cached_tokens_details={"text_tokens": 100, "audio_tokens": 400},
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(model="gpt-realtime-2.1-mini", usage=usage, custom_llm_provider="openai")
|
||||
assert prompt_cost == pytest.approx(300 * 6e-7 + 100 * 6e-8 + 200 * 1e-5 + 400 * 3e-7)
|
||||
|
||||
|
||||
def test_cache_read_breakdown_splits_cached_audio_at_the_audio_cache_rate(_local_model_cost_map: None) -> None:
|
||||
usage = Usage(
|
||||
prompt_tokens=4863,
|
||||
completion_tokens=1087,
|
||||
total_tokens=5950,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=1693,
|
||||
audio_tokens=3170,
|
||||
cached_tokens=2816,
|
||||
cached_tokens_details={"text_tokens": 896, "audio_tokens": 1920},
|
||||
),
|
||||
)
|
||||
|
||||
breakdown = get_token_type_cost_breakdown(model="gpt-realtime-2.1-mini", custom_llm_provider="openai", usage=usage)
|
||||
prompt_cost, _ = generic_cost_per_token(model="gpt-realtime-2.1-mini", usage=usage, custom_llm_provider="openai")
|
||||
|
||||
assert breakdown.cache_read_cost == pytest.approx(896 * 6e-8 + 1920 * 3e-7)
|
||||
assert breakdown.rates is not None
|
||||
assert breakdown.rates.cache_read_input_audio_token_cost == pytest.approx(3e-7)
|
||||
assert prompt_cost == pytest.approx((1693 - 896) * 6e-7 + (3170 - 1920) * 1e-5 + breakdown.cache_read_cost)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "custom_llm_provider", "expected_prompt_cost"),
|
||||
(
|
||||
pytest.param("azure/gpt-realtime-2025-08-28", "azure", 300 * 4e-6 + 100 * 4e-7 + 200 * 3.2e-5 + 400 * 4e-7, id="azure-gpt-realtime"),
|
||||
pytest.param("azure/gpt-realtime-1.5-2026-02-23", "azure", 300 * 4e-6 + 100 * 4e-7 + 200 * 3.2e-5 + 400 * 4e-7, id="azure-gpt-realtime-1.5"),
|
||||
pytest.param("azure/gpt-realtime-mini", "azure", 300 * 6e-7 + 100 * 6e-8 + 200 * 1e-5 + 400 * 3e-7, id="azure-gpt-realtime-mini"),
|
||||
pytest.param("gpt-realtime-mini", "openai", 300 * 6e-7 + 100 * 6e-8 + 200 * 1e-5 + 400 * 3e-7, id="openai-gpt-realtime-mini"),
|
||||
),
|
||||
)
|
||||
def test_realtime_models_bill_cached_text_and_audio_at_their_cache_read_rates(
|
||||
_local_model_cost_map: None, model: str, custom_llm_provider: str, expected_prompt_cost: float
|
||||
) -> None:
|
||||
usage = Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=0,
|
||||
total_tokens=1000,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=400,
|
||||
audio_tokens=600,
|
||||
cached_tokens=500,
|
||||
cached_tokens_details={"text_tokens": 100, "audio_tokens": 400},
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(model=model, usage=usage, custom_llm_provider=custom_llm_provider)
|
||||
assert prompt_cost == pytest.approx(expected_prompt_cost)
|
||||
|
|
|
|||
|
|
@ -11,12 +11,12 @@ import logging
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.fallback_generalizations import (
|
||||
get_fallback_generalization_rules,
|
||||
match_capability_generalizations,
|
||||
match_fill_missing_generalizations,
|
||||
match_routing_generalization,
|
||||
set_fallback_generalizations,
|
||||
)
|
||||
|
|
@ -116,6 +116,67 @@ def test_capability_union_is_last_wins_in_file_order(restore_generalizations):
|
|||
}
|
||||
|
||||
|
||||
def test_fill_missing_requires_per_rule_opt_in(restore_generalizations):
|
||||
restore_generalizations(
|
||||
[
|
||||
{"name": "base", "pattern": r"^acme-", "model_info": {"supports_reasoning": True}},
|
||||
{
|
||||
"name": "opt-in",
|
||||
"pattern": r"^acme-",
|
||||
"fill_missing_for_providers": ["openai"],
|
||||
"model_info": {"supports_vision": True},
|
||||
},
|
||||
]
|
||||
)
|
||||
assert match_fill_missing_generalizations("acme-1", "openai") == {"supports_vision": True}
|
||||
assert match_fill_missing_generalizations("acme-1", "azure") is None
|
||||
assert match_capability_generalizations("acme-1") == {
|
||||
"supports_reasoning": True,
|
||||
"supports_vision": True,
|
||||
}
|
||||
|
||||
restore_generalizations(
|
||||
[{"name": "base", "pattern": r"^acme-", "model_info": {"supports_reasoning": True}}]
|
||||
)
|
||||
assert match_fill_missing_generalizations("acme-1", "openai") is None
|
||||
|
||||
restore_generalizations(
|
||||
[
|
||||
{
|
||||
"name": "mixed",
|
||||
"pattern": r"^acme-",
|
||||
"fill_missing_for_providers": ["openai"],
|
||||
"model_info": {"litellm_provider": "openai", "supports_vision": True},
|
||||
}
|
||||
]
|
||||
)
|
||||
assert match_fill_missing_generalizations("acme-1", "openai") == {"supports_vision": True}
|
||||
|
||||
restore_generalizations(
|
||||
[
|
||||
{
|
||||
"name": "route",
|
||||
"pattern": r"^acme-",
|
||||
"fill_missing_for_providers": ["openai"],
|
||||
"model_info": {"litellm_provider": "openai"},
|
||||
}
|
||||
]
|
||||
)
|
||||
assert match_fill_missing_generalizations("acme-1", "openai") is None
|
||||
|
||||
restore_generalizations(
|
||||
[
|
||||
{
|
||||
"name": "malformed",
|
||||
"pattern": r"^acme-",
|
||||
"fill_missing_for_providers": "openai",
|
||||
"model_info": {"supports_vision": True},
|
||||
}
|
||||
]
|
||||
)
|
||||
assert match_fill_missing_generalizations("acme-1", "openai") is None
|
||||
|
||||
|
||||
def test_routing_rules_are_excluded_from_capability_results(restore_generalizations):
|
||||
restore_generalizations(
|
||||
[
|
||||
|
|
@ -299,6 +360,76 @@ def test_exact_entry_takes_precedence_over_rule(restore_generalizations):
|
|||
assert info["input_cost_per_token"] != 999.0
|
||||
|
||||
|
||||
def test_exact_entries_fill_only_missing_fields(restore_generalizations, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"model_cost",
|
||||
{
|
||||
**litellm.model_cost,
|
||||
"acme-full": {
|
||||
"input_cost_per_token": 1e-6,
|
||||
"output_cost_per_token": 2e-6,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
"max_tokens": 7,
|
||||
"supports_reasoning": False,
|
||||
},
|
||||
"acme-bare": {
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 4e-6,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
"acme-image": {
|
||||
"input_cost_per_token": 5e-6,
|
||||
"output_cost_per_token": 6e-6,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
},
|
||||
"acme-other": {
|
||||
"input_cost_per_token": 7e-6,
|
||||
"output_cost_per_token": 8e-6,
|
||||
"litellm_provider": "openrouter",
|
||||
"mode": "chat",
|
||||
},
|
||||
},
|
||||
)
|
||||
restore_generalizations(
|
||||
[
|
||||
{
|
||||
"name": "acme-backfill",
|
||||
"pattern": r"^acme-",
|
||||
"fill_missing_for_providers": ["openai"],
|
||||
"model_info": {"supports_reasoning": True, "max_tokens": 5},
|
||||
}
|
||||
]
|
||||
)
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
full = litellm.get_model_info("acme-full", custom_llm_provider="openai")
|
||||
assert full["supports_reasoning"] is False
|
||||
assert full["max_tokens"] == 7
|
||||
|
||||
bare = litellm.get_model_info("acme-bare", custom_llm_provider="openai")
|
||||
assert bare["supports_reasoning"] is True
|
||||
assert bare["max_tokens"] == 5
|
||||
assert bare["input_cost_per_token"] == 3e-6
|
||||
assert bare["key"] == "acme-bare"
|
||||
|
||||
other = litellm.get_model_info("acme-other", custom_llm_provider="openrouter")
|
||||
assert other.get("supports_reasoning") is None
|
||||
|
||||
image = litellm.get_model_info("acme-image", custom_llm_provider="openai")
|
||||
assert image.get("supports_reasoning") is None
|
||||
|
||||
restore_generalizations(
|
||||
[{"name": "acme-backfill", "pattern": r"^acme-", "model_info": {"supports_reasoning": True, "max_tokens": 5}}]
|
||||
)
|
||||
litellm.get_model_info.cache_clear()
|
||||
unflagged = litellm.get_model_info("acme-bare", custom_llm_provider="openai")
|
||||
assert unflagged.get("supports_reasoning") is None
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Shipped rules (bundled cost map)
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
|
@ -415,6 +546,18 @@ def test_shipped_version_boundaries(shipped_cost_map, model, provider, adaptive,
|
|||
assert info.get("supports_mid_conversation_system") is mid_conversation, model
|
||||
|
||||
|
||||
def test_shipped_claude_version_regex_excludes_undelimited_41(shipped_cost_map):
|
||||
unmatched = match_capability_generalizations("github_copilot/claude-opus-41")
|
||||
assert unmatched is None or "supports_adaptive_thinking" not in unmatched
|
||||
assert unmatched is None or "supports_mid_conversation_system" not in unmatched
|
||||
|
||||
for model in ("claude-opus-5", "claude-sonnet-4-8"):
|
||||
matched = match_capability_generalizations(model)
|
||||
assert matched is not None
|
||||
assert matched["supports_adaptive_thinking"] is True
|
||||
assert matched["supports_mid_conversation_system"] is True
|
||||
|
||||
|
||||
def test_shipped_rules_cover_new_families_like_fable_at_5_plus(shipped_cost_map):
|
||||
"""Both version gates accept any claude-<family>- id at major 5 or higher, bare
|
||||
major or major-minor, so a new family shaped like claude-fable-5 gets adaptive
|
||||
|
|
@ -605,6 +748,10 @@ def test_shipped_wandb_rule_loses_to_mapped_non_reasoning_entries(shipped_cost_m
|
|||
assert litellm.supports_reasoning(model=model, custom_llm_provider="wandb") is False, model
|
||||
|
||||
|
||||
def test_shipped_wandb_rule_does_not_fill_missing_mapped_entries(shipped_cost_map):
|
||||
assert match_fill_missing_generalizations("wandb/meta-llama/Llama-3.1-8B-Instruct", "wandb") is None
|
||||
|
||||
|
||||
def test_shipped_wandb_rule_is_anchored_to_the_wandb_namespace(shipped_cost_map):
|
||||
"""``^wandb/`` is anchored, so it cannot leak onto another provider's ids."""
|
||||
assert match_capability_generalizations("wandb/some-new-model") == {"supports_reasoning": True}
|
||||
|
|
@ -722,3 +869,58 @@ def test_shipped_openai_reasoning_rule_skips_non_reasoning_gpt_ids(shipped_cost_
|
|||
def test_shipped_openai_reasoning_rule_loses_to_mapped_entries(shipped_cost_map):
|
||||
assert "gpt-5-search-api" in litellm.model_cost
|
||||
assert litellm.supports_reasoning(model="gpt-5-search-api", custom_llm_provider="openai") is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,provider,expected_supports_reasoning",
|
||||
[
|
||||
("azure/us/o1-2024-12-17", "azure", True),
|
||||
("github_copilot/gpt-5", "github_copilot", None),
|
||||
("openrouter/openai/o1", "openrouter", None),
|
||||
("perplexity/openai/gpt-5.4-mini", "perplexity", None),
|
||||
],
|
||||
)
|
||||
def test_shipped_openai_reasoning_rule_backfills_only_approved_providers(
|
||||
shipped_cost_map, model, provider, expected_supports_reasoning
|
||||
):
|
||||
assert model in litellm.model_cost
|
||||
raw_entry = litellm.model_cost[model]
|
||||
assert "supports_reasoning" not in raw_entry
|
||||
model_without_provider = model.removeprefix(f"{provider}/")
|
||||
info = litellm.get_model_info(model=model_without_provider, custom_llm_provider=provider)
|
||||
assert info.get("supports_reasoning") is expected_supports_reasoning
|
||||
assert info["input_cost_per_token"] == raw_entry.get("input_cost_per_token", 0)
|
||||
|
||||
|
||||
def test_shipped_openai_reasoning_rule_matches_only_openai(shipped_cost_map):
|
||||
assert match_fill_missing_generalizations("gpt-5.4", "openai") == {"supports_reasoning": True}
|
||||
assert match_fill_missing_generalizations("gpt-5.4", "openrouter") is None
|
||||
|
||||
|
||||
def test_shipped_openai_reasoning_rule_skips_non_text_modes(shipped_cost_map):
|
||||
model = "gemini/deep-research-pro-preview-12-2025"
|
||||
assert model in litellm.model_cost
|
||||
raw_entry = litellm.model_cost[model]
|
||||
assert "supports_reasoning" not in raw_entry
|
||||
assert raw_entry["mode"] == "image_generation"
|
||||
|
||||
info = litellm.get_model_info("deep-research-pro-preview-12-2025", custom_llm_provider="gemini")
|
||||
assert info.get("supports_reasoning") is None
|
||||
|
||||
|
||||
def test_shipped_claude_thinking_rules_backfill_only_anthropic(shipped_cost_map):
|
||||
model = "perplexity/anthropic/claude-sonnet-4-6"
|
||||
assert model in litellm.model_cost
|
||||
raw_entry = litellm.model_cost[model]
|
||||
assert "supports_adaptive_thinking" not in raw_entry
|
||||
assert "max_input_tokens" not in raw_entry
|
||||
|
||||
info = litellm.get_model_info(model="anthropic/claude-sonnet-4-6", custom_llm_provider="perplexity")
|
||||
assert info.get("supports_adaptive_thinking") is None
|
||||
assert info.get("supports_legacy_thinking") is None
|
||||
assert info.get("max_input_tokens") is None
|
||||
assert match_fill_missing_generalizations("claude-sonnet-4-6", "anthropic") == {
|
||||
"supports_adaptive_thinking": True,
|
||||
"supports_legacy_thinking": True,
|
||||
}
|
||||
assert match_fill_missing_generalizations("claude-sonnet-4-6", "perplexity") is None
|
||||
|
|
|
|||
|
|
@ -5221,10 +5221,11 @@ def test_handle_anthropic_messages_parsed_response_logging_preserves_fast_mode_s
|
|||
assert getattr(result.usage, "speed", None) == "fast"
|
||||
|
||||
|
||||
def test_logging_init_sets_trace_id():
|
||||
def test_logging_init_sets_trace_id(monkeypatch):
|
||||
"""Logging.__init__() must call set_trace_id with self.litellm_trace_id."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
trace_id_var.set("")
|
||||
|
||||
log_obj = Logging(
|
||||
|
|
@ -5240,7 +5241,7 @@ def test_logging_init_sets_trace_id():
|
|||
assert trace_id_var.get() == log_obj.litellm_trace_id
|
||||
|
||||
|
||||
def test_logging_init_skips_stamping_when_correlation_logging_unsupported():
|
||||
def test_logging_init_skips_stamping_when_correlation_logging_unsupported(monkeypatch):
|
||||
"""supports_correlation_logging=False (what wrapper(), the sync entry
|
||||
point, always passes) must leave trace_id_var/session_id_var completely
|
||||
untouched, even though self.litellm_trace_id/litellm_session_id (the
|
||||
|
|
@ -5248,6 +5249,7 @@ def test_logging_init_skips_stamping_when_correlation_logging_unsupported():
|
|||
usual - only the ambient contextvar stamping is gated."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
|
@ -5271,10 +5273,48 @@ def test_logging_init_skips_stamping_when_correlation_logging_unsupported():
|
|||
assert log_obj.litellm_session_id == "should-not-be-stamped"
|
||||
|
||||
|
||||
def test_logging_init_sets_session_id_when_provided():
|
||||
def test_logging_init_skips_stamping_when_request_correlation_in_logs_disabled(monkeypatch):
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
||||
trace_id_var.set("outer")
|
||||
session_id_var.set("outer-sid")
|
||||
try:
|
||||
with (
|
||||
patch( # test-quality-ok: regression test verifies disabled stamping skips both setters
|
||||
"litellm.litellm_core_utils.litellm_logging.set_trace_id"
|
||||
) as mock_set_trace_id,
|
||||
patch( # test-quality-ok: regression test verifies disabled stamping skips both setters
|
||||
"litellm.litellm_core_utils.litellm_logging.set_session_id"
|
||||
) as mock_set_session_id,
|
||||
):
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="call-disabled",
|
||||
function_id="fn-disabled",
|
||||
kwargs={"litellm_session_id": "disabled-session"},
|
||||
supports_correlation_logging=True,
|
||||
)
|
||||
|
||||
assert trace_id_var.get() == "outer"
|
||||
assert session_id_var.get() == "outer-sid"
|
||||
assert log_obj._own_trace_id == "outer"
|
||||
mock_set_trace_id.assert_not_called()
|
||||
mock_set_session_id.assert_not_called()
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_logging_init_sets_session_id_when_provided(monkeypatch):
|
||||
"""Logging.__init__() must call set_session_id when litellm_session_id is in kwargs."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
session_id_var.set("")
|
||||
|
||||
Logging(
|
||||
|
|
@ -5290,11 +5330,12 @@ def test_logging_init_sets_session_id_when_provided():
|
|||
assert session_id_var.get() == "my-session-99"
|
||||
|
||||
|
||||
def test_logging_init_resets_session_id_to_empty_when_absent():
|
||||
def test_logging_init_resets_session_id_to_empty_when_absent(monkeypatch):
|
||||
"""When no session_id is in kwargs, Logging.__init__() must reset session_id_var to ""
|
||||
so a prior request's session_id does not leak into subsequent log records."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
session_id_var.set("preexisting-sid")
|
||||
|
||||
Logging(
|
||||
|
|
@ -5310,7 +5351,7 @@ def test_logging_init_resets_session_id_to_empty_when_absent():
|
|||
assert session_id_var.get() == ""
|
||||
|
||||
|
||||
def test_restore_correlation_context_resets_to_pre_call_value():
|
||||
def test_restore_correlation_context_resets_to_pre_call_value(monkeypatch):
|
||||
"""_restore_correlation_context() must put trace_id_var/session_id_var back to
|
||||
whatever they were immediately before this Logging instance was constructed.
|
||||
This is the mechanism that prevents a nested call (e.g. a guardrail's own
|
||||
|
|
@ -5318,6 +5359,7 @@ def test_restore_correlation_context_resets_to_pre_call_value():
|
|||
session_id into the outer call's subsequent log lines."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
trace_id_var.set("outer-trace")
|
||||
session_id_var.set("outer-session")
|
||||
try:
|
||||
|
|
@ -5343,7 +5385,7 @@ def test_restore_correlation_context_resets_to_pre_call_value():
|
|||
session_id_var.set("")
|
||||
|
||||
|
||||
def test_restore_correlation_context_safe_to_call_repeatedly():
|
||||
def test_restore_correlation_context_safe_to_call_repeatedly(monkeypatch):
|
||||
"""Calling _restore_correlation_context() more than once must not raise.
|
||||
|
||||
It's deliberately NOT guarded against repeat calls: wrapper()'s finally
|
||||
|
|
@ -5353,6 +5395,7 @@ def test_restore_correlation_context_safe_to_call_repeatedly():
|
|||
the contextvars, so repeat calls are expected, not just tolerated."""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
|
|
@ -5367,8 +5410,40 @@ def test_restore_correlation_context_safe_to_call_repeatedly():
|
|||
log_obj._restore_correlation_context() # must not raise
|
||||
|
||||
|
||||
def test_restore_correlation_context_does_not_resanitize(monkeypatch):
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm._logging import _sanitize_correlation_id
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
trace_id_var.set("outer-trace")
|
||||
session_id_var.set("outer-session")
|
||||
try:
|
||||
log_obj = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=None,
|
||||
litellm_call_id="call-no-resanitize",
|
||||
function_id="fn-no-resanitize",
|
||||
kwargs={"litellm_session_id": "inner-session"},
|
||||
)
|
||||
|
||||
with patch( # test-quality-ok: regression test verifies restore avoids sanitization
|
||||
"litellm._logging._sanitize_correlation_id", wraps=_sanitize_correlation_id
|
||||
) as mock_sanitize:
|
||||
log_obj._restore_correlation_context()
|
||||
|
||||
mock_sanitize.assert_not_called()
|
||||
assert trace_id_var.get() == "outer-trace"
|
||||
assert session_id_var.get() == "outer-session"
|
||||
finally:
|
||||
trace_id_var.set("")
|
||||
session_id_var.set("")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restore_correlation_context_works_across_asyncio_task_boundary():
|
||||
async def test_restore_correlation_context_works_across_asyncio_task_boundary(monkeypatch):
|
||||
"""_restore_correlation_context() must succeed even when it's called from a
|
||||
different asyncio Task than the one Logging.__init__() ran in - exactly what
|
||||
happens on litellm's real async success path, where async_success_handler is
|
||||
|
|
@ -5385,6 +5460,7 @@ async def test_restore_correlation_context_works_across_asyncio_task_boundary():
|
|||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
||||
trace_id_var.set("outer-trace-cross-task")
|
||||
session_id_var.set("outer-session-cross-task")
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -10,9 +10,9 @@ from websockets.exceptions import ConnectionClosed
|
|||
from websockets.frames import Close
|
||||
|
||||
import litellm
|
||||
from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.realtime_streaming import (
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
RealTimeStreaming,
|
||||
client_sent_openai_beta_realtime_header,
|
||||
)
|
||||
|
|
@ -3399,6 +3399,26 @@ async def test_refused_session_does_not_stamp_the_reservation_ownership_marker()
|
|||
assert REALTIME_SESSION_SUCCESS_LOGGED_KEY not in session.logging.model_call_details
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refused_session_stamps_the_failure_ownership_marker():
|
||||
"""LIT-6463: the enqueued failure callback releases the key's max_parallel_requests
|
||||
slot from the logging worker, so a refusal stamps REALTIME_SESSION_FAILURE_LOGGED_KEY.
|
||||
The proxy endpoint reads it to leave the slot to that callback instead of racing it.
|
||||
A session that relayed frames logs a success and must not carry the failure stamp."""
|
||||
upstream_close: Final = ConnectionClosed(Close(1008, _UPSTREAM_REFUSAL), None)
|
||||
refused: Final = _relay_session(_client_ws_that_never_sends(), _backend_ws_closing_with(upstream_close))
|
||||
session_created: Final = json.dumps({"type": "session.created", "session": {"id": "sess_1"}}).encode()
|
||||
relayed: Final = _relay_session(
|
||||
_client_ws_that_never_sends(), _backend_ws_closing_with(session_created, upstream_close)
|
||||
)
|
||||
|
||||
await refused.run()
|
||||
await relayed.run()
|
||||
|
||||
assert refused.logging.model_call_details.get(REALTIME_SESSION_FAILURE_LOGGED_KEY) is True
|
||||
assert REALTIME_SESSION_FAILURE_LOGGED_KEY not in relayed.logging.model_call_details
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transformed_transcription_completion_never_sends_response_create():
|
||||
from typing import Final
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue