mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
chore(mcp): reconcile discovery attribution tests with main
This commit is contained in:
commit
cf35e3d376
81 changed files with 2248 additions and 14604 deletions
|
|
@ -131,6 +131,7 @@ GATEWAY_EXACT_PATHS: frozenset[str] = frozenset(
|
|||
"/redoc",
|
||||
"/test",
|
||||
"/debug/memory/summary",
|
||||
"/api/event_logging/batch",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -61,7 +61,7 @@
|
|||
"/v1/fine-tuning" "/fine-tuning" "/v1/responses" "/responses" "/v1/threads" "/threads"
|
||||
"/v1/assistants" "/assistants" "/v1/vector_stores" "/vector_stores" "/v1/indexes"
|
||||
"/v1/models" "/models" "/openai" "/engines"
|
||||
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a" "/a2a"
|
||||
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a" "/a2a" "/api/event_logging"
|
||||
"/v1/rerank" "/v2/rerank" "/rerank" "/v1/ocr" "/ocr" "/v1/rag" "/rag"
|
||||
"/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search"
|
||||
"/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat"
|
||||
|
|
|
|||
|
|
@ -238,6 +238,9 @@ 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"
|
||||
|
||||
# litellm_params flag on failure logs for requests the proxy rejected before routing to a deployment
|
||||
PROXY_REJECTED_BEFORE_ROUTING_KEY: Final = "proxy_rejected_before_routing"
|
||||
|
||||
# Generic fallback for unknown models
|
||||
DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET: Final = int(
|
||||
os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128)
|
||||
|
|
|
|||
|
|
@ -776,9 +776,15 @@ def _joined_choice(parts: tuple[str, ...]) -> tuple[_Choice, ...]:
|
|||
return (_text_choice("\n\n".join(parts)),) if parts else ()
|
||||
|
||||
|
||||
def _text_completion_choice(choice: Mapping[str, object], text: str) -> Mapping[str, object]:
|
||||
synthesized: Final = _text_choice(text, as_str(choice.get("finish_reason")))
|
||||
merged: Final = (*choice.items(), *synthesized.items())
|
||||
return {k: v for k, v in merged if k != "text"} # mutable-ok: mappers json.dumps and isinstance(dict) it
|
||||
|
||||
|
||||
def _completion_choices(response: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
|
||||
return tuple(
|
||||
_text_choice(text, as_str(choice.get("finish_reason")))
|
||||
_text_completion_choice(choice, text)
|
||||
if "message" not in choice and isinstance(text := choice.get("text"), str)
|
||||
else choice
|
||||
for choice in _dicts(response.get("choices"))
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import PROXY_LLM_PROVIDER_FALLBACK
|
||||
from litellm.constants import PROXY_LLM_PROVIDER_FALLBACK, PROXY_REJECTED_BEFORE_ROUTING_KEY
|
||||
from litellm.exceptions import (
|
||||
validate_rate_limit_category,
|
||||
validate_rate_limit_type,
|
||||
|
|
@ -214,17 +214,22 @@ def _get_proxy_llm_router() -> Router | None:
|
|||
return llm_router
|
||||
|
||||
|
||||
def _bounded_requested_model_label(requested_model: str | None, router_originated: bool = False) -> str | None:
|
||||
def _bounded_requested_model_label(requested_model: object, router_originated: bool = False) -> str | None:
|
||||
"""
|
||||
Bound ``requested_model`` label cardinality: names the router recognizes
|
||||
(model names, deployment ids, aliases, routing groups, team public model
|
||||
names) or matches via a global or team wildcard/pattern route keep their
|
||||
own label value; any other client-supplied string collapses into the
|
||||
single ``other`` bucket. With no proxy router to vouch for the string,
|
||||
client-supplied values collapse to ``other`` while ``router_originated``
|
||||
values (emitted by an SDK ``Router``'s own deployment failure and
|
||||
fallback events, where the proxy router never exists) pass through.
|
||||
single ``other`` bucket, as does any non-string request ``model`` value.
|
||||
With no proxy router to vouch for the string, client-supplied values
|
||||
collapse to ``other`` while ``router_originated`` values (emitted by an
|
||||
SDK ``Router``'s own deployment failure and fallback events, where the
|
||||
proxy router never exists) pass through.
|
||||
"""
|
||||
if requested_model is None:
|
||||
return None
|
||||
if not isinstance(requested_model, str):
|
||||
return UNRECOGNIZED_REQUESTED_MODEL_LABEL
|
||||
if not requested_model:
|
||||
return requested_model
|
||||
llm_router: Final = _get_proxy_llm_router()
|
||||
|
|
@ -2832,7 +2837,7 @@ class PrometheusLogger(CustomLogger):
|
|||
|
||||
# On LiteLLM-side rejects (no deployment picked), route request_kwargs["model"]
|
||||
# into requested_model and leave deployment-scoped labels empty.
|
||||
deployment_selected: Final = bool(model_id)
|
||||
deployment_selected: Final = bool(model_id) and not _litellm_params.get(PROXY_REJECTED_BEFORE_ROUTING_KEY)
|
||||
if deployment_selected:
|
||||
label_litellm_model_name = litellm_model_name
|
||||
label_model_id = model_id
|
||||
|
|
|
|||
|
|
@ -82,7 +82,7 @@ def _rejected_image_fetch(url: str, verdict: SSRFError) -> "litellm.ImageFetchEr
|
|||
verbose_logger.warning("Image fetch of %s rejected before any request went out: %s", url, verdict)
|
||||
return litellm.ImageFetchError(
|
||||
"Error: Unable to fetch image from URL. The proxy could not resolve this host or its URL policy rejected it; "
|
||||
f"an admin can check the proxy log and `user_url_allowed_hosts` in general_settings. url={url}"
|
||||
f"an admin can check the proxy log and `user_url_allowed_hosts` in litellm_settings. url={url}"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -336,7 +336,7 @@ def validate_url(url: str) -> tuple[str, str]:
|
|||
raise SSRFError(
|
||||
f"URL targets a blocked address ({resolved_ip}). "
|
||||
"If this is a legitimate internal service, add the host "
|
||||
"to `user_url_allowed_hosts` in general_settings."
|
||||
"to `user_url_allowed_hosts` in litellm_settings."
|
||||
)
|
||||
|
||||
# For HTTPS with SSL verification enabled, TLS certificate validation
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -986,7 +986,7 @@ def _forwarded_upstream_header_names() -> frozenset[str]:
|
|||
)
|
||||
|
||||
|
||||
def _upstream_credential_headers(header_names: Iterable[str]) -> frozenset[str]:
|
||||
def upstream_credential_headers(header_names: Iterable[str]) -> frozenset[str]:
|
||||
"""Lowercased names of the headers in ``header_names`` that carry an upstream MCP
|
||||
credential rather than request context: the configured client side auth header, any
|
||||
header name a configured server forwards upstream via ``extra_headers``, and the
|
||||
|
|
@ -1038,7 +1038,7 @@ def build_synthetic_mcp_request(
|
|||
custom_key_header: Final = _custom_litellm_key_header_name()
|
||||
excluded: Final = (
|
||||
_SYNTHETIC_REQUEST_EXCLUDED_HEADERS
|
||||
| _upstream_credential_headers(raw_headers.keys() if raw_headers else ())
|
||||
| upstream_credential_headers(raw_headers.keys() if raw_headers else ())
|
||||
| (frozenset({custom_key_header.lower()}) if custom_key_header else frozenset())
|
||||
)
|
||||
forwarded: Final = tuple(
|
||||
|
|
@ -1086,7 +1086,7 @@ def logging_safe_mcp_headers(raw_headers: Mapping[str, str] | None) -> Mapping[s
|
|||
)
|
||||
|
||||
excluded: Final = (
|
||||
_upstream_credential_headers(raw_headers.keys() if raw_headers else ())
|
||||
upstream_credential_headers(raw_headers.keys() if raw_headers else ())
|
||||
| UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS
|
||||
| frozenset({"host"})
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2029,7 +2029,14 @@ async def add_litellm_data_to_request(
|
|||
_headers,
|
||||
allow_client_message_redaction_opt_out=_allow_client_message_redaction_opt_out,
|
||||
)
|
||||
_logging_safe_headers: Final = redact_credential_headers(_headers)
|
||||
from litellm.proxy._experimental.mcp_server.utils import upstream_credential_headers
|
||||
|
||||
_mcp_credential_headers: Final = upstream_credential_headers(_headers)
|
||||
_logging_safe_headers: Final = redact_credential_headers(
|
||||
MappingProxyType(
|
||||
{name: value for name, value in _headers.items() if name.lower() not in _mcp_credential_headers}
|
||||
)
|
||||
)
|
||||
verbose_proxy_logger.debug("Request Headers: %s", _logging_safe_headers)
|
||||
verbose_proxy_logger.debug("Raw Headers: %s", _raw_headers)
|
||||
|
||||
|
|
|
|||
|
|
@ -33,21 +33,20 @@ else:
|
|||
FastAPI = Any
|
||||
|
||||
|
||||
def _deprioritize_script_dir_in_sys_path() -> None:
|
||||
def _drop_script_dir_from_sys_path() -> None:
|
||||
"""Stop ``litellm/proxy`` modules from shadowing installed packages.
|
||||
|
||||
Running this file as a script puts its own directory at ``sys.path[0]``, so
|
||||
``import a2a`` resolves to ``litellm/proxy/a2a`` instead of the ``a2a`` SDK
|
||||
and A2A agent calls fail. The entry is moved to the end rather than dropped,
|
||||
because the sibling-import fallbacks in this module (``from proxy_server
|
||||
import ...``) still need it. No-op under the ``litellm`` console script.
|
||||
and ``proxy_server`` resolves to a second copy of
|
||||
``litellm.proxy.proxy_server``. No-op under the ``litellm`` console script.
|
||||
"""
|
||||
script_dir: Final = os.path.dirname(os.path.abspath(__file__))
|
||||
if sys.path and os.path.abspath(sys.path[0]) == script_dir:
|
||||
sys.path.append(sys.path.pop(0))
|
||||
sys.path.pop(0)
|
||||
|
||||
|
||||
_deprioritize_script_dir_in_sys_path()
|
||||
_drop_script_dir_from_sys_path()
|
||||
sys.path.append(os.getcwd())
|
||||
|
||||
config_filename: Final = "litellm.secrets"
|
||||
|
|
@ -881,7 +880,7 @@ class ProxyInitializationHelpers:
|
|||
default=False,
|
||||
help="Use prisma db push instead of prisma migrate for database schema updates",
|
||||
)
|
||||
@click.option("--local", is_flag=True, default=False, help="for local debugging")
|
||||
@click.option("--local", is_flag=True, default=False, help="no-op, kept for backwards compatibility")
|
||||
@click.option(
|
||||
"--skip_server_startup",
|
||||
is_flag=True,
|
||||
|
|
@ -1058,35 +1057,15 @@ def run_server(
|
|||
return
|
||||
|
||||
args: Final = locals()
|
||||
if local:
|
||||
from proxy_server import (
|
||||
try:
|
||||
from litellm.proxy.proxy_server import (
|
||||
KeyManagementSettings,
|
||||
ProxyConfig,
|
||||
app,
|
||||
save_worker_config,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
from .proxy_server import (
|
||||
KeyManagementSettings,
|
||||
ProxyConfig,
|
||||
app,
|
||||
save_worker_config,
|
||||
)
|
||||
except ModuleNotFoundError as e:
|
||||
raise ModuleNotFoundError(f"Missing dependency {e}. Run `pip install 'litellm[proxy]'`")
|
||||
except ImportError as e:
|
||||
if "litellm[proxy]" in str(e):
|
||||
# user is missing a proxy dependency, ask them to pip install litellm[proxy]
|
||||
raise e
|
||||
else:
|
||||
# this is just a local/relative import error, user git cloned litellm
|
||||
from proxy_server import (
|
||||
KeyManagementSettings,
|
||||
ProxyConfig,
|
||||
app,
|
||||
save_worker_config,
|
||||
)
|
||||
except ModuleNotFoundError as e:
|
||||
raise ModuleNotFoundError(f"Missing dependency {e}. Run `pip install 'litellm[proxy]'`") from e
|
||||
if version is True:
|
||||
ProxyInitializationHelpers._echo_litellm_version()
|
||||
return
|
||||
|
|
@ -1283,6 +1262,7 @@ def run_server(
|
|||
|
||||
if os.getenv("DATABASE_URL", None) is not None or os.getenv("DIRECT_URL", None) is not None:
|
||||
from litellm.proxy.db.db_url_settings import (
|
||||
DISABLE_PREPARED_STATEMENTS_ENV_VAR,
|
||||
add_missing_query_params,
|
||||
idle_lifetime_params,
|
||||
reader_shareable_params,
|
||||
|
|
@ -1305,12 +1285,16 @@ def run_server(
|
|||
sys.exit(1)
|
||||
from litellm.secret_managers.main import get_secret
|
||||
|
||||
env_disable_prepared_statements: Final = token_auth_flag_enabled(
|
||||
os.getenv(DISABLE_PREPARED_STATEMENTS_ENV_VAR), env_var=DISABLE_PREPARED_STATEMENTS_ENV_VAR
|
||||
)
|
||||
disable_prepared_statements: Final = db_disable_prepared_statements or env_disable_prepared_statements
|
||||
connection_url_params: Final = _build_db_connection_url_params(
|
||||
connection_limit=db_connection_pool_limit,
|
||||
pool_timeout=db_connection_timeout,
|
||||
connect_timeout=db_connect_timeout,
|
||||
socket_timeout=db_socket_timeout,
|
||||
disable_prepared_statements=db_disable_prepared_statements,
|
||||
disable_prepared_statements=disable_prepared_statements,
|
||||
extra_params=db_extra_connection_params,
|
||||
)
|
||||
lifetime_params: Final = idle_lifetime_params(general_settings.get("database_max_idle_connection_lifetime"))
|
||||
|
|
|
|||
|
|
@ -52,6 +52,7 @@ from litellm.constants import (
|
|||
DEFAULT_MODEL_CREATED_AT_TIME,
|
||||
LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL,
|
||||
MAX_TEAM_LIST_LIMIT,
|
||||
PROXY_REJECTED_BEFORE_ROUTING_KEY,
|
||||
REDIS_SPEND_LOGS_BUFFER_DEQUEUE_COUNT,
|
||||
SPEND_LOG_QUEUE_MAX_BYTES,
|
||||
SPEND_LOG_WRITE_BATCH_MAX_BYTES,
|
||||
|
|
@ -979,6 +980,92 @@ def _failure_usage_to_lift(
|
|||
_EMPTY_LIFT: Final = MappingProxyType({})
|
||||
|
||||
|
||||
def _stamp_deployment_attribution(
|
||||
litellm_params: dict[str, object], model_group: str | None, team_id: str | None, dispatched: bool
|
||||
) -> Mapping[str, object]:
|
||||
"""Stamp provider and logging-metadata attribution onto ``litellm_params`` and return it.
|
||||
``litellm_params["model_info"]`` stays unset: the router's cooldown and per-deployment rpm
|
||||
callbacks key off it and must not count a proxy-side reject against the deployment. A failure
|
||||
after the provider handoff keeps the metadata the router stamped; a request that never reached a
|
||||
provider is flagged ``PROXY_REJECTED_BEFORE_ROUTING_KEY`` (deployment metrics key off it) whatever
|
||||
its metadata says, since ``metadata.model_info`` can be caller supplied."""
|
||||
attribution: Final = _deployment_attribution_for_model_group(model_group, team_id)
|
||||
if "custom_llm_provider" in attribution:
|
||||
litellm_params["custom_llm_provider"] = attribution["custom_llm_provider"]
|
||||
if dispatched:
|
||||
return attribution
|
||||
litellm_params[PROXY_REJECTED_BEFORE_ROUTING_KEY] = True
|
||||
if "model_info" not in attribution:
|
||||
return attribution
|
||||
if litellm_params.get("metadata") is None:
|
||||
litellm_params["metadata"] = {} # mutable-ok: legacy logging payload is populated in place
|
||||
metadata: Final = litellm_params["metadata"]
|
||||
if not isinstance(metadata, dict):
|
||||
return attribution
|
||||
metadata.setdefault("model_info", attribution["model_info"])
|
||||
metadata.setdefault("deployment", attribution["deployment"])
|
||||
return attribution
|
||||
|
||||
|
||||
def _deployment_attribution_for_model_group(model_group: object, team_id: str | None) -> Mapping[str, object]:
|
||||
"""Provider fields the router would have stamped had it reached a deployment:
|
||||
``custom_llm_provider`` when every deployment in the group resolves to the same
|
||||
provider, plus ``model_info`` and ``deployment`` when the group has exactly one.
|
||||
``team_id`` picks the key's team deployments over a global group of the same public name."""
|
||||
if not isinstance(model_group, str):
|
||||
return _EMPTY_LIFT
|
||||
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is None:
|
||||
return _EMPTY_LIFT
|
||||
deployments: Final = llm_router.get_model_list(model_name=model_group, team_id=team_id)
|
||||
if not deployments:
|
||||
return _EMPTY_LIFT
|
||||
|
||||
def _provider_for_deployment(deployment: Mapping[str, object]) -> str | None:
|
||||
litellm_params: Final = cast( # cast-ok: router deployment parameters are mapping-shaped
|
||||
Mapping[str, object], deployment["litellm_params"]
|
||||
)
|
||||
try:
|
||||
provider: Final = litellm.get_llm_provider(
|
||||
model=cast(str, litellm_params["model"]), # cast-ok: router deployment model is a string
|
||||
custom_llm_provider=cast( # cast-ok: router deployment provider is optional
|
||||
str | None, litellm_params.get("custom_llm_provider")
|
||||
),
|
||||
)[1]
|
||||
return cast(str | None, provider) # cast-ok: provider resolver returns an optional provider string
|
||||
except Exception: # noqa: BLE001 # get_llm_provider raises for unmapped models
|
||||
return None
|
||||
|
||||
providers: Final = frozenset(_provider_for_deployment(deployment) for deployment in deployments)
|
||||
shared_provider: Final = next(iter(providers)) if len(providers) == 1 else None
|
||||
single_deployment: Final = deployments[0] if len(deployments) == 1 else None
|
||||
single_deployment_params: Final = (
|
||||
cast( # cast-ok: router deployment parameters are mapping-shaped
|
||||
Mapping[str, object], single_deployment["litellm_params"]
|
||||
)
|
||||
if single_deployment is not None
|
||||
else None
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
# mutable-ok: frozen immediately by the outer MappingProxyType
|
||||
**({"custom_llm_provider": shared_provider} if shared_provider is not None else {}),
|
||||
**(
|
||||
{ # mutable-ok: frozen immediately by the outer MappingProxyType
|
||||
"model_info": dict( # mutable-ok: preserve the router's mutable model-info payload
|
||||
single_deployment.get("model_info") or {}
|
||||
),
|
||||
"deployment": single_deployment_params["model"],
|
||||
}
|
||||
if single_deployment is not None and single_deployment_params is not None
|
||||
else {} # mutable-ok: frozen immediately by the outer MappingProxyType
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _call_type_for_route(route: str | None) -> str | None:
|
||||
"""The route's call type when it maps to a single operation (its async and sync variants);
|
||||
None for routes shared by several operations, since the method is not known here."""
|
||||
|
|
@ -3227,11 +3314,25 @@ class ProxyLogging:
|
|||
elif k not in ("model", "user", "litellm_logging_obj"):
|
||||
_optional_params[k] = v
|
||||
|
||||
attribution: Final = _stamp_deployment_attribution(
|
||||
_litellm_params,
|
||||
request_data.get("model"),
|
||||
user_api_key_dict.team_id,
|
||||
dispatched=litellm_logging_obj.model_call_details.get("first_api_call_start_time") is not None,
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model=request_data.get("model", ""),
|
||||
user=request_data.get("user", ""),
|
||||
optional_params=_optional_params,
|
||||
litellm_params=_litellm_params,
|
||||
**(
|
||||
{ # mutable-ok: frozen immediately by keyword expansion
|
||||
"custom_llm_provider": attribution["custom_llm_provider"]
|
||||
}
|
||||
if "custom_llm_provider" in attribution
|
||||
else {} # mutable-ok: frozen immediately by keyword expansion
|
||||
),
|
||||
)
|
||||
|
||||
input: list | str | dict = ""
|
||||
|
|
|
|||
|
|
@ -209,16 +209,12 @@ async def aresponses_api_with_mcp(
|
|||
user_api_key_auth = kwargs.get("user_api_key_auth") or kwargs.get("litellm_metadata", {}).get("user_api_key_auth")
|
||||
|
||||
# Extract MCP auth headers from request (for dynamic auth when fetching tools)
|
||||
mcp_auth_header: str | None = None
|
||||
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None
|
||||
secret_fields = kwargs.get("secret_fields")
|
||||
if secret_fields and isinstance(secret_fields, dict):
|
||||
(
|
||||
mcp_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
_,
|
||||
_,
|
||||
) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request(secret_fields=secret_fields, tools=tools)
|
||||
secret_fields: Final = kwargs.get("secret_fields")
|
||||
mcp_auth_header, mcp_server_auth_headers, _, discovery_raw_headers = (
|
||||
ResponsesAPIRequestUtils.extract_mcp_headers_from_request(secret_fields=secret_fields, tools=tools)
|
||||
if isinstance(secret_fields, dict) and secret_fields
|
||||
else (None, None, None, None)
|
||||
)
|
||||
|
||||
# Get original MCP tools (for events) and OpenAI tools (for LLM) by reusing existing methods
|
||||
(
|
||||
|
|
@ -231,6 +227,7 @@ async def aresponses_api_with_mcp(
|
|||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs),
|
||||
raw_headers=discovery_raw_headers,
|
||||
)
|
||||
openai_tools: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(original_mcp_tools)
|
||||
|
||||
|
|
@ -330,7 +327,6 @@ async def aresponses_api_with_mcp(
|
|||
user_api_key_auth = kwargs.get("litellm_metadata", {}).get("user_api_key_auth")
|
||||
|
||||
# Extract MCP auth headers from the request to pass to MCP server
|
||||
secret_fields = kwargs.get("secret_fields")
|
||||
(
|
||||
mcp_auth_header,
|
||||
mcp_server_auth_headers,
|
||||
|
|
@ -416,6 +412,7 @@ async def aresponses_api_with_mcp(
|
|||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs),
|
||||
raw_headers=discovery_raw_headers,
|
||||
)
|
||||
final_response = LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response(
|
||||
response=final_response,
|
||||
|
|
|
|||
|
|
@ -138,6 +138,7 @@ async def acompletion_with_mcp(
|
|||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request_tags=request_tags,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
openai_tools: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(
|
||||
|
|
|
|||
|
|
@ -239,6 +239,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
mcp_auth_header: str | None = None,
|
||||
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list[MCPTool], list[str]]:
|
||||
"""
|
||||
Get available tools from the MCP server manager.
|
||||
|
|
@ -329,6 +330,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
list_tools_log_source="responses",
|
||||
litellm_trace_id=litellm_trace_id,
|
||||
request_tags=request_tags,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
tools: Final = listing.tools
|
||||
|
||||
|
|
@ -455,6 +457,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
mcp_auth_header: str | None = None,
|
||||
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
|
||||
request_tags: list[str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
) -> tuple[list[MCPTool], dict[str, str]]:
|
||||
"""
|
||||
Process MCP tools through filtering and deduplication pipeline without OpenAI transformation.
|
||||
|
|
@ -485,6 +488,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
request_tags=request_tags,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
# Step 2: Filter tools based on allowed_tools parameter
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -74,7 +74,7 @@ locals {
|
|||
"/v1/models*", "/models*",
|
||||
"/openai/*", "/engines/*",
|
||||
"/v1/messages*", "/messages*",
|
||||
"/v1/skills/*", "/v1/a2a/*",
|
||||
"/v1/skills/*", "/v1/a2a/*", "/api/event_logging*",
|
||||
"/v1/rerank*", "/v2/rerank*", "/rerank*",
|
||||
"/v1/ocr*", "/ocr*",
|
||||
"/v1/rag/*", "/rag/*",
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ locals {
|
|||
"/v1/models*", "/models*",
|
||||
"/openai/*", "/engines/*",
|
||||
"/v1/messages*", "/messages*",
|
||||
"/v1/skills/*", "/v1/a2a/*",
|
||||
"/v1/skills/*", "/v1/a2a/*", "/api/event_logging*",
|
||||
"/v1/rerank*", "/v2/rerank*", "/rerank*",
|
||||
"/v1/ocr*", "/ocr*",
|
||||
"/v1/rag/*", "/rag/*",
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from __future__ import annotations
|
|||
import ast
|
||||
import atexit
|
||||
import hashlib
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
|
|
@ -15,10 +16,18 @@ import socket
|
|||
import sys
|
||||
import threading
|
||||
from collections import defaultdict
|
||||
from typing import Iterable
|
||||
from collections.abc import Iterable, Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest import mock
|
||||
|
||||
import aiohttp
|
||||
import pytest
|
||||
import vcr
|
||||
import vcr.matchers as _vcr_matchers
|
||||
import vcr.patch as _vcr_patch
|
||||
|
||||
from tests._vcr_redis_persister import (
|
||||
MAX_EPISODES_PER_CASSETTE,
|
||||
|
|
@ -127,9 +136,7 @@ def emit_vcr_diagnostic_log(terminalreporter) -> None:
|
|||
with open(path, "r", encoding="utf-8") as fh:
|
||||
content = fh.read()
|
||||
except OSError as exc:
|
||||
read_errors.append(
|
||||
f" [failed to read {name}: {type(exc).__name__}: {exc}]"
|
||||
)
|
||||
read_errors.append(f" [failed to read {name}: {type(exc).__name__}: {exc}]")
|
||||
continue
|
||||
for line in content.splitlines():
|
||||
if not line.strip():
|
||||
|
|
@ -142,9 +149,7 @@ def emit_vcr_diagnostic_log(terminalreporter) -> None:
|
|||
return
|
||||
|
||||
terminalreporter.write_sep("=", "VCR DIAGNOSTIC LOG", bold=True)
|
||||
terminalreporter.write_line(
|
||||
f" source dir: {directory} (deduplicated; full log archived as a CI artifact)"
|
||||
)
|
||||
terminalreporter.write_line(f" source dir: {directory} (deduplicated; full log archived as a CI artifact)")
|
||||
for line in read_errors:
|
||||
terminalreporter.write_line(line)
|
||||
|
||||
|
|
@ -235,9 +240,7 @@ def pin_httpx_multipart_boundary(monkeypatch) -> None:
|
|||
boundary = VCR_FIXED_MULTIPART_BOUNDARY.encode("ascii")
|
||||
return _original_init(self, data=data, files=files, boundary=boundary, **kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
_httpx_multipart.MultipartStream, "__init__", _init_with_fixed_boundary
|
||||
)
|
||||
monkeypatch.setattr(_httpx_multipart.MultipartStream, "__init__", _init_with_fixed_boundary)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
|
|
@ -270,11 +273,7 @@ def _replace_b64_json_in_place(obj) -> bool:
|
|||
changed = False
|
||||
if isinstance(obj, dict):
|
||||
for key, value in obj.items():
|
||||
if (
|
||||
key == "b64_json"
|
||||
and isinstance(value, str)
|
||||
and len(value) > len(VCR_IMAGE_B64_PLACEHOLDER)
|
||||
):
|
||||
if key == "b64_json" and isinstance(value, str) and len(value) > len(VCR_IMAGE_B64_PLACEHOLDER):
|
||||
obj[key] = VCR_IMAGE_B64_PLACEHOLDER
|
||||
changed = True
|
||||
elif _replace_b64_json_in_place(value):
|
||||
|
|
@ -296,16 +295,12 @@ def _strip_image_b64_payloads(response):
|
|||
preserves all those checks while shrinking cassettes by ~99%.
|
||||
"""
|
||||
if not isinstance(response, dict):
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-strip-b64] response is {type(response).__name__!r}, not "
|
||||
"dict; skipping b64 scrub"
|
||||
)
|
||||
vcr_diag_write_line(f"[vcr-strip-b64] response is {type(response).__name__!r}, not dict; skipping b64 scrub")
|
||||
return response
|
||||
body = response.get("body")
|
||||
if not isinstance(body, dict):
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-strip-b64] response['body'] is {type(body).__name__!r}, "
|
||||
"not dict; skipping b64 scrub"
|
||||
f"[vcr-strip-b64] response['body'] is {type(body).__name__!r}, not dict; skipping b64 scrub"
|
||||
)
|
||||
return response
|
||||
raw = body.get("string")
|
||||
|
|
@ -316,10 +311,7 @@ def _strip_image_b64_payloads(response):
|
|||
try:
|
||||
text = bytes(raw).decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
vcr_diag_write_line(
|
||||
"[vcr-strip-b64] response body bytes are not valid UTF-8; "
|
||||
"skipping b64 scrub"
|
||||
)
|
||||
vcr_diag_write_line("[vcr-strip-b64] response body bytes are not valid UTF-8; skipping b64 scrub")
|
||||
return response
|
||||
was_bytes = True
|
||||
elif isinstance(raw, str):
|
||||
|
|
@ -327,8 +319,7 @@ def _strip_image_b64_payloads(response):
|
|||
was_bytes = False
|
||||
else:
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-strip-b64] response['body']['string'] is "
|
||||
f"{type(raw).__name__!r}, not bytes/str; skipping b64 scrub"
|
||||
f"[vcr-strip-b64] response['body']['string'] is {type(raw).__name__!r}, not bytes/str; skipping b64 scrub"
|
||||
)
|
||||
return response
|
||||
|
||||
|
|
@ -349,9 +340,7 @@ def _strip_image_b64_payloads(response):
|
|||
for key in list(headers):
|
||||
if str(key).lower() == "content-length":
|
||||
value = headers[key]
|
||||
headers[key] = (
|
||||
[new_len_value] if isinstance(value, list) else new_len_value
|
||||
)
|
||||
headers[key] = [new_len_value] if isinstance(value, list) else new_len_value
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -409,15 +398,11 @@ def _canonical_body(request) -> tuple[bytes, str]:
|
|||
# selected. This mirrors the existing SigV4 / multipart-boundary / b64-image
|
||||
# normalizations already in this module, and means the already-bloated
|
||||
# cassettes start replaying immediately without a flush + re-record.
|
||||
_VCR_UUID_RE = re.compile(
|
||||
rb"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}"
|
||||
)
|
||||
_VCR_UUID_RE = re.compile(rb"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}")
|
||||
_VCR_LITELLM_BATCH_JOB_RE = re.compile(rb"litellm-batch-[0-9a-fA-F]{8}")
|
||||
# ISO-8601 timestamps, e.g. ``2026-05-25T03:40:37.262045Z`` /
|
||||
# ``2026-05-25T03:40:37+00:00``.
|
||||
_VCR_ISO_TS_RE = re.compile(
|
||||
rb"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?"
|
||||
)
|
||||
_VCR_ISO_TS_RE = re.compile(rb"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?")
|
||||
# Unix epoch as 13-digit milliseconds, then 10-digit ``time.time()`` float,
|
||||
# then 10-digit integer seconds. Anchored to ``1`` + 9/12 digits, which keeps
|
||||
# them inside the 2001-2033 / 2001-2033 epoch windows and avoids matching
|
||||
|
|
@ -639,10 +624,7 @@ def _should_drop_telemetry_record(request) -> bool:
|
|||
return False
|
||||
if not _is_telemetry_request(request):
|
||||
return False
|
||||
if (
|
||||
_is_telemetry_export_request(request)
|
||||
and not _current_test_replays_telemetry_export()
|
||||
):
|
||||
if _is_telemetry_export_request(request) and not _current_test_replays_telemetry_export():
|
||||
return True
|
||||
return not _current_test_records_telemetry()
|
||||
|
||||
|
|
@ -767,9 +749,7 @@ def _iter_header_values(headers, name: str):
|
|||
yield value
|
||||
|
||||
|
||||
_AWS_SIGV4_CREDENTIAL_RE = re.compile(
|
||||
r"AWS4-HMAC-SHA256\s+Credential=([^/\s,]+)/", re.IGNORECASE
|
||||
)
|
||||
_AWS_SIGV4_CREDENTIAL_RE = re.compile(r"AWS4-HMAC-SHA256\s+Credential=([^/\s,]+)/", re.IGNORECASE)
|
||||
|
||||
# Google OAuth2 access tokens always start with ``ya29.`` regardless of how
|
||||
# they were minted (service account, metadata server, impersonation).
|
||||
|
|
@ -891,9 +871,7 @@ def _normalize_multipart_boundary(request) -> None:
|
|||
return
|
||||
|
||||
try:
|
||||
headers[content_type_key] = content_type_value.replace(
|
||||
match.group(0), fixed_param
|
||||
)
|
||||
headers[content_type_key] = content_type_value.replace(match.group(0), fixed_param)
|
||||
except (TypeError, AttributeError):
|
||||
return
|
||||
|
||||
|
|
@ -985,8 +963,7 @@ def _materialize_iterable_body(request) -> None:
|
|||
uri = getattr(request, "uri", getattr(request, "url", "?"))
|
||||
first_type = type(chunks[0]).__name__ if chunks else "empty"
|
||||
vcr_diag_write_line(
|
||||
f"[vcr-materialize] FALLBACK: {method} {uri} chunk type "
|
||||
f"{first_type!r} not coerced to bytes; storing b''"
|
||||
f"[vcr-materialize] FALLBACK: {method} {uri} chunk type {first_type!r} not coerced to bytes; storing b''"
|
||||
)
|
||||
out = b""
|
||||
|
||||
|
|
@ -1026,9 +1003,7 @@ def _key_fingerprint_matcher(r1, r2) -> None:
|
|||
return
|
||||
|
||||
def _fp(req):
|
||||
for value in _iter_header_values(
|
||||
getattr(req, "headers", None), KEY_FINGERPRINT_HEADER
|
||||
):
|
||||
for value in _iter_header_values(getattr(req, "headers", None), KEY_FINGERPRINT_HEADER):
|
||||
if value is None:
|
||||
continue
|
||||
return value if isinstance(value, str) else str(value)
|
||||
|
|
@ -1159,13 +1134,11 @@ def _print_atexit_banner() -> None:
|
|||
_emit("VCR CASSETTE CACHE DEGRADED")
|
||||
if save_failures:
|
||||
_emit(
|
||||
f" {save_failures} cassette save failure(s); last error: "
|
||||
f"{health.get('save_failure_last_error', '')}"
|
||||
f" {save_failures} cassette save failure(s); last error: {health.get('save_failure_last_error', '')}"
|
||||
)
|
||||
if load_failures:
|
||||
_emit(
|
||||
f" {load_failures} cassette load failure(s); last error: "
|
||||
f"{health.get('load_failure_last_error', '')}"
|
||||
f" {load_failures} cassette load failure(s); last error: {health.get('load_failure_last_error', '')}"
|
||||
)
|
||||
if snapshot:
|
||||
_emit(_format_capacity_line(snapshot))
|
||||
|
|
@ -1276,11 +1249,7 @@ class _RespxUsageVisitor(ast.NodeVisitor):
|
|||
if isinstance(dec, ast.Call):
|
||||
dec = dec.func
|
||||
if isinstance(dec, ast.Attribute):
|
||||
return (
|
||||
isinstance(dec.value, ast.Name)
|
||||
and dec.value.id == "respx"
|
||||
and dec.attr == "mock"
|
||||
)
|
||||
return isinstance(dec.value, ast.Name) and dec.value.id == "respx" and dec.attr == "mock"
|
||||
return False
|
||||
|
||||
def _is_pytest_mark_respx(self, dec: ast.expr) -> bool:
|
||||
|
|
@ -1307,9 +1276,7 @@ class _RespxUsageVisitor(ast.NodeVisitor):
|
|||
# ``def test_foo(respx_mock): ...`` — pytest supplies the fixture
|
||||
# whenever the parameter name appears, regardless of marker.
|
||||
all_args = (
|
||||
list(args.args)
|
||||
+ list(args.kwonlyargs)
|
||||
+ (list(args.posonlyargs) if hasattr(args, "posonlyargs") else [])
|
||||
list(args.args) + list(args.kwonlyargs) + (list(args.posonlyargs) if hasattr(args, "posonlyargs") else [])
|
||||
)
|
||||
for a in all_args:
|
||||
if a.arg == "respx_mock":
|
||||
|
|
@ -1566,9 +1533,7 @@ def _emit_outcome_payload(
|
|||
},
|
||||
)
|
||||
)
|
||||
node.user_properties.append(
|
||||
(_USER_PROP_RECORDED_BY, os.environ.get("PYTEST_XDIST_WORKER", ""))
|
||||
)
|
||||
node.user_properties.append((_USER_PROP_RECORDED_BY, os.environ.get("PYTEST_XDIST_WORKER", "")))
|
||||
|
||||
|
||||
def aggregate_report_outcome(report) -> None:
|
||||
|
|
@ -1616,9 +1581,7 @@ def aggregate_report_outcome(report) -> None:
|
|||
if verdict == VERDICT_MISS_OVERFLOW:
|
||||
_session_stats["overflow_tests"].append(nodeid)
|
||||
elif verdict == VERDICT_UNMARKED_LIVE_CALL:
|
||||
_session_stats["unmarked_live_call_tests"].append(
|
||||
(nodeid, list(outcome.get("live_call_hosts") or []))
|
||||
)
|
||||
_session_stats["unmarked_live_call_tests"].append((nodeid, list(outcome.get("live_call_hosts") or [])))
|
||||
|
||||
skip_reason = outcome.get("skip_reason")
|
||||
if skip_reason:
|
||||
|
|
@ -1635,9 +1598,7 @@ def session_stats_snapshot() -> dict:
|
|||
"overflow_tests": list(_session_stats["overflow_tests"]),
|
||||
"unmarked_live_call_tests": list(_session_stats["unmarked_live_call_tests"]),
|
||||
"skip_reason_counts": dict(_session_stats["skip_reason_counts"]),
|
||||
"skip_reason_examples": {
|
||||
k: list(v) for k, v in _session_stats["skip_reason_examples"].items()
|
||||
},
|
||||
"skip_reason_examples": {k: list(v) for k, v in _session_stats["skip_reason_examples"].items()},
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -1810,9 +1771,7 @@ def record_vcr_outcome(request, vcr) -> None:
|
|||
# Cassette is None ⇒ test wasn't VCR-marked. Honor the skip reason
|
||||
# we tagged at collection time, and pull live-call hosts captured by
|
||||
# the socket probe (if any).
|
||||
skip_reason = getattr(
|
||||
request.node, VCR_SKIP_REASON_USER_ATTR, SKIP_REASON_FILE_OPT_OUT
|
||||
)
|
||||
skip_reason = getattr(request.node, VCR_SKIP_REASON_USER_ATTR, SKIP_REASON_FILE_OPT_OUT)
|
||||
_session_stats["skip_reason_counts"][skip_reason] += 1
|
||||
|
||||
hosts = getattr(request.node, _LIVE_CALL_BUFFER_KEY, []) or []
|
||||
|
|
@ -1837,9 +1796,7 @@ def record_vcr_outcome(request, vcr) -> None:
|
|||
live_call_hosts=hosts,
|
||||
)
|
||||
if vcr_outcome_logging_enabled():
|
||||
request.node.user_properties.append(
|
||||
(_USER_PROP_VERDICT_LINE, _format_verdict_line(verdict, None, extra))
|
||||
)
|
||||
request.node.user_properties.append((_USER_PROP_VERDICT_LINE, _format_verdict_line(verdict, None, extra)))
|
||||
|
||||
|
||||
def install_live_call_probe(request, vcr) -> None:
|
||||
|
|
@ -1858,9 +1815,7 @@ def install_live_call_probe(request, vcr) -> None:
|
|||
# Track the current test for telemetry-leak suppression (applies to every
|
||||
# test, VCR-marked or not). See ``_should_drop_telemetry_record``.
|
||||
global _current_test_nodeid
|
||||
_current_test_nodeid = str(
|
||||
getattr(getattr(request, "node", None), "nodeid", "") or ""
|
||||
)
|
||||
_current_test_nodeid = str(getattr(getattr(request, "node", None), "nodeid", "") or "")
|
||||
if vcr is not None or vcr_disabled():
|
||||
return None
|
||||
probe = _LiveCallProbe()
|
||||
|
|
@ -1876,10 +1831,7 @@ def _format_capacity_line(snapshot: dict) -> str:
|
|||
pct = float(snapshot.get("used_pct", 0.0) or 0.0)
|
||||
used_mb = used / (1024 * 1024)
|
||||
cap_mb = cap / (1024 * 1024)
|
||||
return (
|
||||
f" Cassette Redis usage: {used_mb:.1f} MiB / {cap_mb:.1f} MiB "
|
||||
f"({pct:.1f}% of maxmemory)"
|
||||
)
|
||||
return f" Cassette Redis usage: {used_mb:.1f} MiB / {cap_mb:.1f} MiB ({pct:.1f}% of maxmemory)"
|
||||
|
||||
|
||||
def emit_vcr_classification_summary(terminalreporter) -> None:
|
||||
|
|
@ -1940,14 +1892,10 @@ def emit_vcr_classification_summary(terminalreporter) -> None:
|
|||
total_leaks = sum(leak_counts.values())
|
||||
terminalreporter.write_sep("-", "VCR COST LEAK CHECK", bold=True)
|
||||
if total_leaks:
|
||||
rendered = ", ".join(
|
||||
f"{verdict}={count}" for verdict, count in leak_counts.items() if count
|
||||
)
|
||||
rendered = ", ".join(f"{verdict}={count}" for verdict, count in leak_counts.items() if count)
|
||||
terminalreporter.write_line(f" FAIL: {rendered}")
|
||||
else:
|
||||
terminalreporter.write_line(
|
||||
" PASS: no overflow, partial, not-persisted, or unmarked live-call verdicts"
|
||||
)
|
||||
terminalreporter.write_line(" PASS: no overflow, partial, not-persisted, or unmarked live-call verdicts")
|
||||
|
||||
overflow = snapshot["overflow_tests"]
|
||||
if overflow:
|
||||
|
|
@ -2007,18 +1955,14 @@ def emit_cassette_cache_session_banner(terminalreporter) -> None:
|
|||
snapshot = cassette_cache_capacity_snapshot()
|
||||
|
||||
if save_failures or load_failures:
|
||||
terminalreporter.write_sep(
|
||||
"=", "VCR CASSETTE CACHE DEGRADED", red=True, bold=True
|
||||
)
|
||||
terminalreporter.write_sep("=", "VCR CASSETTE CACHE DEGRADED", red=True, bold=True)
|
||||
if save_failures:
|
||||
terminalreporter.write_line(
|
||||
f" {save_failures} cassette save failure(s); last error: "
|
||||
f"{health.get('save_failure_last_error', '')}"
|
||||
f" {save_failures} cassette save failure(s); last error: {health.get('save_failure_last_error', '')}"
|
||||
)
|
||||
if load_failures:
|
||||
terminalreporter.write_line(
|
||||
f" {load_failures} cassette load failure(s); last error: "
|
||||
f"{health.get('load_failure_last_error', '')}"
|
||||
f" {load_failures} cassette load failure(s); last error: {health.get('load_failure_last_error', '')}"
|
||||
)
|
||||
terminalreporter.write_line(
|
||||
" Tests still passed because cassette persistence is best-effort, "
|
||||
|
|
@ -2031,9 +1975,7 @@ def emit_cassette_cache_session_banner(terminalreporter) -> None:
|
|||
return
|
||||
|
||||
if snapshot and snapshot["used_pct"] >= CASSETTE_CACHE_HIGH_WATER_FRACTION * 100:
|
||||
terminalreporter.write_sep(
|
||||
"=", "VCR CASSETTE CACHE NEAR CAPACITY", yellow=True, bold=True
|
||||
)
|
||||
terminalreporter.write_sep("=", "VCR CASSETTE CACHE NEAR CAPACITY", yellow=True, bold=True)
|
||||
terminalreporter.write_line(_format_capacity_line(snapshot))
|
||||
terminalreporter.write_line(
|
||||
" No save failures yet, but Redis is approaching maxmemory. "
|
||||
|
|
@ -2082,13 +2024,104 @@ class VerboseReporterState:
|
|||
if reporter is None:
|
||||
return
|
||||
verdict = next(
|
||||
(
|
||||
v
|
||||
for k, v in (report.user_properties or [])
|
||||
if k == _USER_PROP_VERDICT_LINE
|
||||
),
|
||||
(v for k, v in (report.user_properties or []) if k == _USER_PROP_VERDICT_LINE),
|
||||
None,
|
||||
)
|
||||
if not verdict:
|
||||
return
|
||||
reporter.write_line(f"{verdict} :: {report.nodeid}")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VcrPatchPoint:
|
||||
owner: object
|
||||
attribute: str
|
||||
original: object
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return f"{_patch_owner_name(self.owner)}.{self.attribute}"
|
||||
|
||||
def current(self) -> object:
|
||||
current: Final[object] = getattr(self.owner, self.attribute)
|
||||
return current
|
||||
|
||||
def is_patched(self) -> bool:
|
||||
return self.current() is not self.original
|
||||
|
||||
def restore(self) -> None:
|
||||
setattr(self.owner, self.attribute, self.original)
|
||||
|
||||
|
||||
def _patch_owner_name(owner: object) -> str:
|
||||
if inspect.isclass(owner):
|
||||
return f"{owner.__module__}.{owner.__qualname__}"
|
||||
if inspect.ismodule(owner):
|
||||
return owner.__name__
|
||||
return repr(owner)
|
||||
|
||||
|
||||
def _vcr_patch_point(patcher: mock._patch[object]) -> VcrPatchPoint:
|
||||
owner: Final[object] = patcher.getter()
|
||||
return VcrPatchPoint(owner=owner, attribute=patcher.attribute, original=patcher.new)
|
||||
|
||||
|
||||
_VCR_PATCH_POINTS: Final = (
|
||||
*(_vcr_patch_point(patcher) for patcher in _vcr_patch.reset_patchers()),
|
||||
VcrPatchPoint(aiohttp.ClientSession, "_request", _vcr_patch._AiohttpClientSessionRequest),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VcrPatchLeak:
|
||||
patch_points: tuple[str, ...]
|
||||
cassette_paths: tuple[str, ...]
|
||||
|
||||
|
||||
def _cassette_paths_wrapped_into(fn: object) -> tuple[str, ...]:
|
||||
if not inspect.isfunction(fn):
|
||||
return ()
|
||||
cassette: Final = inspect.getclosurevars(fn).nonlocals.get("cassette")
|
||||
own: Final = (str(cassette._path),) if isinstance(cassette, vcr.cassette.Cassette) else ()
|
||||
return own + _cassette_paths_wrapped_into(getattr(fn, "__wrapped__", None))
|
||||
|
||||
|
||||
def detect_vcr_patch_leak() -> VcrPatchLeak | None:
|
||||
leaked: Final = tuple(point for point in _VCR_PATCH_POINTS if point.is_patched())
|
||||
if not leaked:
|
||||
return None
|
||||
return VcrPatchLeak(
|
||||
patch_points=tuple(point.name for point in leaked),
|
||||
cassette_paths=tuple(
|
||||
dict.fromkeys(path for point in leaked for path in _cassette_paths_wrapped_into(point.current()))
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def restore_vcr_patch_points() -> None:
|
||||
for point in _VCR_PATCH_POINTS:
|
||||
point.restore()
|
||||
|
||||
|
||||
def guard_vcr_patch_points(item: pytest.Item, teardown_failed: bool) -> None:
|
||||
leak: Final = detect_vcr_patch_leak()
|
||||
if leak is None:
|
||||
return
|
||||
restore_vcr_patch_points()
|
||||
if teardown_failed:
|
||||
return
|
||||
pytest.fail(
|
||||
f"{item.nodeid} finished with a vcrpy cassette still patched into "
|
||||
f"{', '.join(leak.patch_points)} (cassettes: {', '.join(leak.cassette_paths) or 'unknown'}); "
|
||||
"the originals were restored so later tests are unaffected",
|
||||
pytrace=False,
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def rewound_new_episodes_cassette(cassette_dir: Path) -> Iterator[vcr.cassette.Cassette]:
|
||||
cassette_path: Final = cassette_dir / "rewound_owner.yaml"
|
||||
cassette_path.write_text("interactions: []\nversion: 1\n")
|
||||
recorder: Final = vcr.VCR(cassette_library_dir=str(cassette_dir))
|
||||
with recorder.use_cassette(cassette_path.name, record_mode="new_episodes") as cassette:
|
||||
yield cassette
|
||||
|
|
|
|||
25
tests/capturing_transport.py
Normal file
25
tests/capturing_transport.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
class CapturingTransport(httpx.AsyncBaseTransport, httpx.BaseTransport):
|
||||
def __init__(self, response: BaseModel) -> None:
|
||||
self._response: Final = response
|
||||
self.request_bodies: tuple[Mapping[str, object], ...] = ()
|
||||
|
||||
def handle_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return self._respond(request.read())
|
||||
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return self._respond(await request.aread())
|
||||
|
||||
def _respond(self, body: bytes) -> httpx.Response:
|
||||
self.request_bodies = (*self.request_bodies, _JSON_OBJECT.validate_json(body))
|
||||
return httpx.Response(200, json=self._response.model_dump(mode="json"))
|
||||
|
|
@ -216,7 +216,7 @@ quota_management.<behavior>.<variant>.<assertion>
|
|||
| isolates_per_model | isolates_per_member | isolates_per_group | enforced_across_keys
|
||||
| routes_to_fallback | reseed_matches_db | reports_spend | logs_cost | zero_cost
|
||||
| matches_sum_of_logs | loses_no_spend | attributes_spend | writes_own_rows
|
||||
| writes_failure_row | returns_cost | keeps_total | joins_key | reports_alias_and_email
|
||||
| writes_failure_row | attributes_provider | returns_cost | keeps_total | joins_key | reports_alias_and_email
|
||||
| health_rows_keep_service_account | retrieve_batch_cost_joins_retrieving_key
|
||||
| poller_batch_cost_joins_creating_key | bills_under_request_session
|
||||
e.g. quota_management.ratelimit.rpm.blocks_over_limit exercised_on=[chat_completions, messages]
|
||||
|
|
|
|||
|
|
@ -51,6 +51,7 @@
|
|||
- {id: quota_management.spend_tracking.end_user.attributes_spend, module: quota_management, tier: P1, behavior: spend_tracking, variant: end_user, assertions: [attributes_spend], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "user= attribution lands the end-user id on the spend row"}
|
||||
- {id: quota_management.spend_tracking.per_model.writes_own_rows, module: quota_management, tier: P2, behavior: spend_tracking, variant: per_model, assertions: [writes_own_rows], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "Each model on a shared key gets its own spend row"}
|
||||
- {id: quota_management.spend_tracking.failure.writes_failure_row, module: quota_management, tier: P1, behavior: spend_tracking, variant: failure, assertions: [writes_failure_row], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_log_error_logger.py", rationale: "A failed call writes a failure-status spend row"}
|
||||
- {id: quota_management.spend_tracking.failure.attributes_provider, module: quota_management, tier: P1, behavior: spend_tracking, variant: failure, assertions: [attributes_provider], exercised_on: [chat_completions], source: "proxy/utils.py", rationale: "A request rejected in pre_call_hook (rate limit, guardrail) still lands its single deployment's provider and model_id on the failure spend row"}
|
||||
- {id: quota_management.spend_tracking.spend_calculate.returns_cost, module: quota_management, tier: P2, behavior: spend_tracking, variant: spend_calculate, assertions: [returns_cost], exercised_on: [spend_calculate], source: "proxy/spend_tracking/spend_management_endpoints.py", rationale: "/spend/calculate prices a hypothetical request at nonzero cost"}
|
||||
- {id: quota_management.spend_tracking.pagination.keeps_total, module: quota_management, tier: P2, behavior: spend_tracking, variant: pagination, assertions: [keeps_total], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_management_endpoints.py", rationale: "Spend-logs v2 pagination caps page size without losing the total"}
|
||||
- {id: quota_management.spend_tracking.cache_write.bills_cache_creation_rate, module: quota_management, tier: P1, behavior: spend_tracking, variant: cache_write, assertions: [bills_cache_creation_rate], exercised_on: [chat_completions], source: "litellm_core_utils/llm_cost_calc/utils.py", rationale: "OpenAI cache-write tokens land on the spend row as cache-creation tokens billed at the cache-creation rate, not silently at the input rate (#34046)"}
|
||||
|
|
|
|||
|
|
@ -952,6 +952,7 @@ class SpendLogRow(BaseModel):
|
|||
cache_hit: str | None = None
|
||||
call_type: str | None = None
|
||||
custom_llm_provider: str | None = None
|
||||
model_id: str | None = None
|
||||
team_id: str | None = None
|
||||
user: str | None = None
|
||||
end_user: str | None = None
|
||||
|
|
|
|||
|
|
@ -21,9 +21,9 @@ from math import isclose
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from e2e_http import Success
|
||||
from e2e_http import RateLimitedError, Success
|
||||
from lifecycle import ResourceManager
|
||||
from models import LiteLLMParamsBody, SpendLogs, SpendLogsParams
|
||||
from models import KeyGenerateBody, LiteLLMParamsBody, SpendLogs, SpendLogsParams
|
||||
from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
|
@ -43,6 +43,7 @@ def _summarize(rows: list[SpendLogRow]) -> list[dict[str, object]]:
|
|||
"cache_hit",
|
||||
"call_type",
|
||||
"custom_llm_provider",
|
||||
"model_id",
|
||||
"prompt_tokens",
|
||||
"completion_tokens",
|
||||
"total_tokens",
|
||||
|
|
@ -498,6 +499,47 @@ def test_failure_call_writes_failure_status_row(
|
|||
assert (failure_row.spend or 0) == 0.0, "failed call must not be charged"
|
||||
|
||||
|
||||
@pytest.mark.covers("quota_management.spend_tracking.failure.attributes_provider")
|
||||
def test_pre_call_rejection_row_attributes_provider_and_model_id(
|
||||
client: SpendClient, resources: ResourceManager
|
||||
) -> None:
|
||||
"""A request the proxy rejects before the router picks a deployment (here the
|
||||
key's rpm limit, a pre_call_hook 429) never reaches the code that stamps the
|
||||
deployment onto the log. The failure row must still carry the provider and
|
||||
model_id of the model group's only deployment, so per-provider failure reports
|
||||
can count it."""
|
||||
model = f"e2e-spend-precall-{unique_marker()}"
|
||||
model_id = client.proxy.create_model(
|
||||
model, LiteLLMParamsBody(model="openai/gpt-5.5", api_key="os.environ/OPENAI_API_KEY")
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
key = client.proxy.generate_key(KeyGenerateBody(models=[model], rpm_limit=1))
|
||||
resources.defer(lambda: client.proxy.delete_key(key))
|
||||
|
||||
unwrap(client.chat(key, model, f"reply with one word {unique_marker()}", max_tokens=8))
|
||||
rejected = client.chat(key, model, f"over the rpm limit {unique_marker()}", max_tokens=8)
|
||||
assert isinstance(rejected, RateLimitedError), (
|
||||
f"the second call on an rpm_limit=1 key must be rejected with 429 before routing, got {rejected}"
|
||||
)
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
key,
|
||||
min_rows=2,
|
||||
predicate=lambda rs: {r.status for r in rs} >= {"success", "failure"},
|
||||
)
|
||||
success_row = _require_row(rows, lambda r: r.status == "success", "for the served call")
|
||||
failure_row = _require_row(rows, lambda r: r.status == "failure", "for the rate-limited call")
|
||||
|
||||
assert failure_row.custom_llm_provider == success_row.custom_llm_provider, (
|
||||
f"rejected call lost its provider: failure row {failure_row.custom_llm_provider!r} vs "
|
||||
f"served row {success_row.custom_llm_provider!r}; {_summarize(rows)}"
|
||||
)
|
||||
assert failure_row.model_id == model_id, (
|
||||
f"rejected call lost its deployment: failure row model_id {failure_row.model_id!r} vs "
|
||||
f"registered {model_id!r}; {_summarize(rows)}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.covers("quota_management.spend_tracking.spend_calculate.returns_cost")
|
||||
def test_spend_calculate_returns_nonzero_cost(client: SpendClient) -> None:
|
||||
cost = client.calculate_spend(
|
||||
|
|
|
|||
|
|
@ -248,6 +248,9 @@
|
|||
"other.observability.callbacks.credentials_stay_out_of_event_bodies",
|
||||
"other.observability.callbacks.concurrent_results_join_complete_events_and_rows"
|
||||
],
|
||||
"tests/integration/observability/test_otel_text_completion_choices.py::test_otel_weave_output_keeps_text_completion_provider_fields_beside_the_synthesized_message": [
|
||||
"other.observability.otel.text_completion_choices_keep_provider_fields"
|
||||
],
|
||||
"tests/integration/observability/test_guardrail_effects.py::test_guardrail_rewrites_system_and_user_in_actual_anthropic_request": [
|
||||
"other.observability.guardrails.rewrite_reaches_correct_anthropic_positions"
|
||||
],
|
||||
|
|
|
|||
|
|
@ -0,0 +1,110 @@
|
|||
import json
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
|
||||
|
||||
def _span_attributes(body: bytes) -> tuple[dict[str, object], ...]:
|
||||
return tuple(
|
||||
{attribute["key"]: attribute["value"] for attribute in span.get("attributes", ())}
|
||||
for resource in json.loads(body)["resourceSpans"]
|
||||
for scope in resource["scopeSpans"]
|
||||
for span in scope["spans"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.covers("other.observability.otel.text_completion_choices_keep_provider_fields")
|
||||
def test_otel_weave_output_keeps_text_completion_provider_fields_beside_the_synthesized_message(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
marker: Final = "otel-text-" + uuid.uuid4().hex
|
||||
logprobs: Final = {
|
||||
"tokens": ["Hello", " there"],
|
||||
"token_logprobs": [-0.1, -0.2],
|
||||
"top_logprobs": None,
|
||||
"text_offset": [0, 5],
|
||||
}
|
||||
content_filter: Final = {"hate": {"filtered": False, "severity": "safe"}}
|
||||
|
||||
def upstream(request: Request) -> Reply:
|
||||
assert request.target.endswith("/completions"), request.target
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": marker,
|
||||
"object": "text_completion",
|
||||
"created": 1,
|
||||
"model": "gpt-3.5-turbo-instruct",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"text": "Hello there",
|
||||
"finish_reason": "stop",
|
||||
"logprobs": logprobs,
|
||||
"content_filter_results": content_filter,
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
def sink(_request: Request) -> Reply:
|
||||
return Reply()
|
||||
|
||||
with wire_server(upstream) as provider, wire_server(sink) as collector:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["litellm_settings"].update({"callbacks": ["otel"]})
|
||||
config["callback_settings"] = {
|
||||
"otel": {
|
||||
"exporter": "http/json",
|
||||
"endpoint": collector.url,
|
||||
"mapper_names": ["genai", "openinference", "weave"],
|
||||
"capture_message_content": "span_only",
|
||||
"use_simple_processor": True,
|
||||
}
|
||||
}
|
||||
path: Final = tmp_path / "otel.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {"LITELLM_OTEL_V2": "1"}, config=path) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(model="openai/gpt-3.5-turbo-instruct", api_base=provider.url + "/v1")
|
||||
response: Final = candidate.request(
|
||||
"POST", "/v1/completions", {"model": model, "prompt": marker, "cache": {"no-cache": True}}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["choices"][0]["text"] == "Hello there"
|
||||
batches = []
|
||||
|
||||
def outputs() -> tuple[list[dict[str, object]], ...]:
|
||||
batches.extend(collector.drain())
|
||||
return tuple(
|
||||
json.loads(attributes["weave.output"]["stringValue"])
|
||||
for batch in batches
|
||||
for attributes in _span_attributes(batch.body)
|
||||
if "weave.output" in attributes
|
||||
and attributes.get("gen_ai.response.id", {}).get("stringValue") == marker
|
||||
)
|
||||
|
||||
choices: Final = eventually(outputs, lambda values: len(values) == 1, seconds=20)[0]
|
||||
assert len(choices) == 1, choices
|
||||
choice: Final = choices[0]
|
||||
assert choice["message"]["content"] == "Hello there", choice
|
||||
assert "text" not in choice, choice
|
||||
assert {
|
||||
key: choice.get(key) for key in ("index", "finish_reason", "logprobs", "content_filter_results")
|
||||
} == {
|
||||
"index": 0,
|
||||
"finish_reason": "stop",
|
||||
"logprobs": logprobs,
|
||||
"content_filter_results": content_filter,
|
||||
}, choice
|
||||
|
|
@ -7,6 +7,8 @@
|
|||
|
||||
import asyncio
|
||||
import importlib
|
||||
from collections.abc import Generator
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -20,6 +22,7 @@ from tests._vcr_conftest_common import ( # noqa: E402,F401
|
|||
emit_cassette_cache_session_banner,
|
||||
emit_vcr_classification_summary,
|
||||
emit_vcr_diagnostic_log,
|
||||
guard_vcr_patch_points,
|
||||
install_live_call_probe,
|
||||
record_vcr_outcome,
|
||||
register_persister_if_enabled,
|
||||
|
|
@ -37,17 +40,12 @@ def fake_openai_endpoint():
|
|||
|
||||
# Per-item respx detection (``apply_vcr_auto_marker_to_items``) handles
|
||||
# the vast majority of respx-vs-vcrpy conflicts automatically. The entries
|
||||
# below are the persister's and the WebSocket VCR's own unit-test files, which
|
||||
# exercise ``save_cassette`` / ``load_cassette`` against fakeredis and must not
|
||||
# themselves run under a live cassette context.
|
||||
# below are the persister's, the WebSocket VCR's, and the cassette patch-leak
|
||||
# guard's own unit-test files, which exercise ``save_cassette`` /
|
||||
# ``load_cassette`` against fakeredis or enter cassettes themselves and must
|
||||
# not run under a live cassette context.
|
||||
_VCR_AUTO_MARKER_SKIP_FILES = frozenset(
|
||||
{"test_vcr_redis_persister.py", "test_ws_vcr.py"}
|
||||
)
|
||||
|
||||
_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = (
|
||||
"test_nvidia_nim.py::test_embedding_nvidia_nim",
|
||||
"test_litellm_proxy_provider.py::test_litellm_gateway_from_sdk_embedding[False]",
|
||||
"test_litellm_proxy_provider.py::test_litellm_gateway_from_sdk_embedding[True]",
|
||||
{"test_vcr_redis_persister.py", "test_ws_vcr.py", "test_vcr_leak_guard.py"}
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -77,6 +75,17 @@ def _vcr_outcome_gate(request, vcr):
|
|||
record_vcr_outcome(request, vcr)
|
||||
|
||||
|
||||
@pytest.hookimpl(wrapper=True, trylast=True)
|
||||
def pytest_runtest_teardown(item: pytest.Item) -> Generator[None, object, object]:
|
||||
try:
|
||||
result: Final = yield
|
||||
except BaseException:
|
||||
guard_vcr_patch_points(item, teardown_failed=True)
|
||||
raise
|
||||
guard_vcr_patch_points(item, teardown_failed=False)
|
||||
return result
|
||||
|
||||
|
||||
def pytest_configure(config):
|
||||
_verbose_state.remember_pluginmanager(config)
|
||||
reset_vcr_diag_dir()
|
||||
|
|
@ -172,7 +181,6 @@ def pytest_collection_modifyitems(config, items):
|
|||
apply_vcr_auto_marker_to_items(
|
||||
items,
|
||||
skip_files=_VCR_AUTO_MARKER_SKIP_FILES,
|
||||
skip_nodeid_suffixes=_VCR_INCOMPATIBLE_NODEID_SUFFIXES,
|
||||
)
|
||||
|
||||
custom_logger_tests = [
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@ import json
|
|||
import re
|
||||
from datetime import datetime
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
||||
|
|
@ -12,7 +14,12 @@ import pytest
|
|||
from unittest.mock import MagicMock, patch
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
|
||||
import pytest_asyncio
|
||||
from openai import AsyncOpenAI
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
from openai.types import CreateEmbeddingResponse, Embedding
|
||||
from openai.types.create_embedding_response import Usage
|
||||
|
||||
from tests.capturing_transport import CapturingTransport
|
||||
from tests._vcr_conftest_common import rewound_new_episodes_cassette
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -87,62 +94,61 @@ async def test_litellm_gateway_from_sdk_structured_output():
|
|||
assert "json_schema" in json_schema
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
_GATEWAY_EMBEDDING_RESPONSE: Final = CreateEmbeddingResponse(
|
||||
object="list",
|
||||
data=(Embedding(object="embedding", index=0, embedding=(0.1, 0.2, 0.3)),),
|
||||
model="my-vllm-model",
|
||||
usage=Usage(prompt_tokens=2, total_tokens=2),
|
||||
)
|
||||
|
||||
|
||||
async def _gateway_embedding_via_injected_client(
|
||||
is_async: bool,
|
||||
) -> tuple[CapturingTransport, litellm.EmbeddingResponse]:
|
||||
transport: Final = CapturingTransport(_GATEWAY_EMBEDDING_RESPONSE)
|
||||
response: Final = (
|
||||
await litellm.aembedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=AsyncOpenAI(api_key="fake-key", http_client=httpx.AsyncClient(transport=transport)),
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
if is_async
|
||||
else litellm.embedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=OpenAI(api_key="fake-key", http_client=httpx.Client(transport=transport)),
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
)
|
||||
return transport, response
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", (False, True))
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_embedding(is_async):
|
||||
async def test_litellm_gateway_from_sdk_embedding(is_async: bool):
|
||||
litellm.set_verbose = True
|
||||
litellm._turn_on_debug()
|
||||
|
||||
captured_bodies = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_bodies.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
|
||||
"model": "my-vllm-model",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
|
||||
if is_async:
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
openai_client = AsyncOpenAI(
|
||||
api_key="fake-key",
|
||||
http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
response = await litellm.aembedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
else:
|
||||
from openai import OpenAI
|
||||
|
||||
openai_client = OpenAI(
|
||||
api_key="fake-key",
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
response = litellm.embedding(
|
||||
model="litellm_proxy/my-vllm-model",
|
||||
input="Hello world",
|
||||
client=openai_client,
|
||||
api_base="my-custom-api-base",
|
||||
)
|
||||
|
||||
request_body = captured_bodies[0]
|
||||
print("Request body - {}".format(request_body))
|
||||
transport, response = await _gateway_embedding_via_injected_client(is_async)
|
||||
|
||||
request_body: Final = transport.request_bodies[0]
|
||||
assert "Hello world" == request_body["input"]
|
||||
assert "my-vllm-model" == request_body["model"]
|
||||
assert "encoding_format" not in request_body
|
||||
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_embedding_under_foreign_cassette(tmp_path: Path):
|
||||
with rewound_new_episodes_cassette(tmp_path):
|
||||
sync_transport, _ = await _gateway_embedding_via_injected_client(is_async=False)
|
||||
async_transport, _ = await _gateway_embedding_via_injected_client(is_async=True)
|
||||
|
||||
assert tuple(body["input"] for body in sync_transport.request_bodies) == ("Hello world",)
|
||||
assert tuple(body["input"] for body in async_transport.request_bodies) == ("Hello world",)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_async", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_litellm_gateway_from_sdk_image_generation(is_async):
|
||||
|
|
|
|||
|
|
@ -1,17 +1,21 @@
|
|||
import json
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai.types import CreateEmbeddingResponse, Embedding
|
||||
from openai.types.create_embedding_response import Usage as EmbeddingUsage
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
import litellm
|
||||
from litellm import Choices, Message, ModelResponse, EmbeddingResponse, Usage
|
||||
from litellm import completion
|
||||
from base_rerank_unit_tests import BaseLLMRerankTest
|
||||
from tests.capturing_transport import CapturingTransport
|
||||
|
||||
|
||||
def test_completion_nvidia_nim():
|
||||
|
|
@ -63,33 +67,23 @@ def test_embedding_nvidia_nim():
|
|||
litellm.set_verbose = True
|
||||
from openai import OpenAI
|
||||
|
||||
captured_bodies = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_bodies.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
|
||||
"model": "nvidia/nv-embedqa-e5-v5",
|
||||
"usage": {"prompt_tokens": 6, "total_tokens": 6},
|
||||
},
|
||||
transport: Final = CapturingTransport(
|
||||
CreateEmbeddingResponse(
|
||||
object="list",
|
||||
data=(Embedding(object="embedding", index=0, embedding=(0.1, 0.2, 0.3)),),
|
||||
model="nvidia/nv-embedqa-e5-v5",
|
||||
usage=EmbeddingUsage(prompt_tokens=6, total_tokens=6),
|
||||
)
|
||||
|
||||
client = OpenAI(
|
||||
api_key="fake-api-key",
|
||||
http_client=httpx.Client(transport=httpx.MockTransport(handler)),
|
||||
)
|
||||
response = litellm.embedding(
|
||||
client: Final = OpenAI(api_key="fake-api-key", http_client=httpx.Client(transport=transport))
|
||||
response: Final = litellm.embedding(
|
||||
model="nvidia_nim/nvidia/nv-embedqa-e5-v5",
|
||||
input="What is the meaning of life?",
|
||||
input_type="passage",
|
||||
dimensions=1024,
|
||||
client=client,
|
||||
)
|
||||
request_body = captured_bodies[0]
|
||||
print("request_body: ", request_body)
|
||||
request_body: Final = transport.request_bodies[0]
|
||||
assert request_body["input"] == "What is the meaning of life?"
|
||||
assert request_body["model"] == "nvidia/nv-embedqa-e5-v5"
|
||||
assert request_body["input_type"] == "passage"
|
||||
|
|
|
|||
72
tests/llm_translation/test_vcr_leak_guard.py
Normal file
72
tests/llm_translation/test_vcr_leak_guard.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import httpx2
|
||||
import pytest
|
||||
|
||||
from tests._vcr_conftest_common import (
|
||||
detect_vcr_patch_leak,
|
||||
guard_vcr_patch_points,
|
||||
restore_vcr_patch_points,
|
||||
rewound_new_episodes_cassette,
|
||||
)
|
||||
|
||||
_ORIGINAL_MOCK_HANDLE_ASYNC_REQUEST: Final = httpx.MockTransport.handle_async_request
|
||||
_ORIGINAL_HTTPX2_MOCK_HANDLE_ASYNC_REQUEST: Final = httpx2.MockTransport.handle_async_request
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def leaked_cassette_dir(tmp_path: Path):
|
||||
context: Final = rewound_new_episodes_cassette(tmp_path)
|
||||
context.__enter__()
|
||||
yield tmp_path
|
||||
context.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_no_leak_when_no_cassette_is_active():
|
||||
assert detect_vcr_patch_leak() is None
|
||||
|
||||
|
||||
def test_leaked_cassette_is_detected_named_and_restorable(leaked_cassette_dir: Path):
|
||||
leak: Final = detect_vcr_patch_leak()
|
||||
|
||||
assert leak is not None
|
||||
assert {"httpx.MockTransport.handle_async_request", "aiohttp.client.ClientSession._request"} <= set(
|
||||
leak.patch_points
|
||||
)
|
||||
assert leak.cassette_paths == (str(leaked_cassette_dir / "rewound_owner.yaml"),)
|
||||
|
||||
restore_vcr_patch_points()
|
||||
|
||||
assert detect_vcr_patch_leak() is None
|
||||
assert httpx.MockTransport.handle_async_request is _ORIGINAL_MOCK_HANDLE_ASYNC_REQUEST
|
||||
|
||||
|
||||
def test_leak_is_detected_on_every_transport_family_vcrpy_patches(leaked_cassette_dir: Path):
|
||||
leak: Final = detect_vcr_patch_leak()
|
||||
|
||||
assert leak is not None
|
||||
assert "httpx2.MockTransport.handle_async_request" in leak.patch_points
|
||||
assert httpx2.MockTransport.handle_async_request is not _ORIGINAL_HTTPX2_MOCK_HANDLE_ASYNC_REQUEST
|
||||
|
||||
restore_vcr_patch_points()
|
||||
|
||||
assert httpx2.MockTransport.handle_async_request is _ORIGINAL_HTTPX2_MOCK_HANDLE_ASYNC_REQUEST
|
||||
|
||||
|
||||
def test_guard_fails_the_leaking_test_and_restores_the_originals(request, leaked_cassette_dir: Path):
|
||||
with pytest.raises(pytest.fail.Exception, match=re.escape(request.node.nodeid)) as failure:
|
||||
guard_vcr_patch_points(request.node, teardown_failed=False)
|
||||
|
||||
assert str(leaked_cassette_dir / "rewound_owner.yaml") in str(failure.value)
|
||||
assert detect_vcr_patch_leak() is None
|
||||
|
||||
|
||||
def test_guard_restores_silently_when_the_teardown_already_failed(request, leaked_cassette_dir: Path):
|
||||
guard_vcr_patch_points(request.node, teardown_failed=True)
|
||||
|
||||
assert detect_vcr_patch_leak() is None
|
||||
|
|
@ -93,7 +93,6 @@ _VCR_INCOMPATIBLE_FILES = frozenset(
|
|||
# carry no real provider cost.
|
||||
_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = (
|
||||
"test_router.py::test_router_text_completion_client",
|
||||
"test_embedding.py::test_encoding_format_omitted_by_default_for_openai_sdk",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -15,6 +15,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import litellm
|
||||
from litellm import completion, completion_cost, embedding
|
||||
from openai.types import CreateEmbeddingResponse
|
||||
from openai.types.create_embedding_response import Usage as EmbeddingUsage
|
||||
from tests.capturing_transport import CapturingTransport
|
||||
from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE
|
||||
|
||||
litellm.set_verbose = False
|
||||
|
|
@ -1268,23 +1271,15 @@ def test_encoding_format_omitted_by_default_for_openai_sdk(monkeypatch):
|
|||
Optional global override: `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT`.
|
||||
"""
|
||||
monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False)
|
||||
captured_bodies = []
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured_bodies.append(json.loads(request.content))
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
|
||||
"model": "text-embedding-ada-002",
|
||||
"usage": {"prompt_tokens": 1, "total_tokens": 1},
|
||||
},
|
||||
transport = CapturingTransport(
|
||||
CreateEmbeddingResponse(
|
||||
object="list",
|
||||
data=(Embedding(object="embedding", index=0, embedding=(0.1, 0.2, 0.3)),),
|
||||
model="text-embedding-ada-002",
|
||||
usage=EmbeddingUsage(prompt_tokens=1, total_tokens=1),
|
||||
)
|
||||
|
||||
client = openai.OpenAI(
|
||||
api_key="sk-test", http_client=httpx.Client(transport=httpx.MockTransport(handler))
|
||||
)
|
||||
client = openai.OpenAI(api_key="sk-test", http_client=httpx.Client(transport=transport))
|
||||
|
||||
response = embedding(
|
||||
model="text-embedding-ada-002",
|
||||
|
|
@ -1294,7 +1289,7 @@ def test_encoding_format_omitted_by_default_for_openai_sdk(monkeypatch):
|
|||
)
|
||||
|
||||
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
assert "encoding_format" not in captured_bodies[0], (
|
||||
assert "encoding_format" not in transport.request_bodies[0], (
|
||||
"encoding_format should be omitted from the upstream request when not provided by user"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -934,11 +934,32 @@ def test_text_completion_choices_become_assistant_messages_in_choice_order() ->
|
|||
capture_content=True,
|
||||
)
|
||||
|
||||
assert data.choices_out == (_assistant_choice(" first", "length"), _assistant_choice(" second", "stop"))
|
||||
assert data.choices_out == (
|
||||
{"index": 0, "logprobs": None, **_assistant_choice(" first", "length")},
|
||||
{"index": 1, "logprobs": None, **_assistant_choice(" second", "stop")},
|
||||
)
|
||||
assert data.finish_reasons == ("length", "stop")
|
||||
assert data.response_id == "cmpl-1"
|
||||
|
||||
|
||||
def test_text_completion_choices_keep_provider_fields_beside_the_synthesized_message() -> None:
|
||||
choice: Final = {
|
||||
"index": 2,
|
||||
"text": "Hello there",
|
||||
"finish_reason": "stop",
|
||||
"logprobs": {"tokens": ["Hello"], "token_logprobs": [-0.1]},
|
||||
"content_filter_results": {"hate": {"filtered": False}},
|
||||
"provider_specific": {"cached": True},
|
||||
}
|
||||
data: Final = LLMCallSpanData.from_standard_logging_payload(
|
||||
_route_payload("atext_completion", "gpt-3.5-turbo-instruct", {"choices": [choice]}), capture_content=True
|
||||
)
|
||||
|
||||
assert data.choices_out == (
|
||||
{k: v for k, v in choice.items() if k != "text"} | _assistant_choice("Hello there", "stop"),
|
||||
)
|
||||
|
||||
|
||||
def test_text_completion_choices_follow_the_content_capture_gate_but_finish_reasons_do_not() -> None:
|
||||
data: Final = LLMCallSpanData.from_standard_logging_payload(
|
||||
_route_payload(
|
||||
|
|
|
|||
|
|
@ -1664,7 +1664,7 @@ class TestEnableAnthropicPromptCaching:
|
|||
points = self._points(model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", provider="bedrock")
|
||||
assert [p["index"] for p in points] == [None, -1]
|
||||
|
||||
@pytest.mark.parametrize("model, provider", [("gpt-4o", "openai"), ("gemini-2.0-flash", "gemini")])
|
||||
@pytest.mark.parametrize("model, provider", [("gpt-4o", "openai")])
|
||||
def test_non_anthropic_providers_never_injected(self, monkeypatch, model, provider):
|
||||
"""These report supports_prompt_caching=True but never consume cache_control markers."""
|
||||
from litellm.utils import supports_prompt_caching
|
||||
|
|
|
|||
|
|
@ -0,0 +1,101 @@
|
|||
"""
|
||||
LIT-7701 attributes a pre_call_hook rejection's failure log to the model
|
||||
group's single deployment (``model_id`` and ``custom_llm_provider``) and flags
|
||||
it with ``PROXY_REJECTED_BEFORE_ROUTING_KEY``. The deployment health metrics
|
||||
must keep treating such rejects as "no deployment picked": a key rate limit or
|
||||
guardrail block never reached the deployment, so it must not flip
|
||||
``litellm_deployment_state`` to partial outage or count as a deployment failure
|
||||
response. A failure raised after the router picked a deployment (a post-call
|
||||
guardrail block, a provider error) carries no flag and keeps its deployment labels.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from prometheus_client import REGISTRY
|
||||
|
||||
from litellm.constants import PROXY_REJECTED_BEFORE_ROUTING_KEY
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def cleanup_prometheus_registry():
|
||||
for collector in list(REGISTRY._collector_to_names.keys()):
|
||||
try:
|
||||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
yield
|
||||
|
||||
for collector in list(REGISTRY._collector_to_names.keys()):
|
||||
try:
|
||||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _attributed_failure_kwargs(exception: Exception, rejected_before_routing: bool) -> dict:
|
||||
return {
|
||||
"model": "openai/gpt-4.1",
|
||||
"litellm_params": {
|
||||
"custom_llm_provider": "openai",
|
||||
"metadata": {"model_info": {"id": "dep-1"}, "model_group": "internal-model"},
|
||||
**({PROXY_REJECTED_BEFORE_ROUTING_KEY: True} if rejected_before_routing else {}),
|
||||
},
|
||||
"standard_logging_object": {
|
||||
"model_id": "dep-1",
|
||||
"model_group": "internal-model",
|
||||
"api_base": "https://api.openai.com",
|
||||
"metadata": {},
|
||||
},
|
||||
"exception": exception,
|
||||
}
|
||||
|
||||
|
||||
def _model_id_values(metric) -> set[str]:
|
||||
index = metric._labelnames.index("model_id")
|
||||
return {sample_key[index] for sample_key in metric._metrics}
|
||||
|
||||
|
||||
class _ProviderError(Exception):
|
||||
status_code = 500
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rejection",
|
||||
[
|
||||
HTTPException(status_code=403, detail="guardrail blocked"),
|
||||
ProxyException(message="budget exceeded", type="budget_exceeded", param=None, code=400),
|
||||
ProxyRateLimitError(detail={"error": "key rpm limit"}),
|
||||
GuardrailRaisedException(guardrail_name="pii", message="blocked", status_code=403),
|
||||
],
|
||||
ids=["http_exception", "proxy_exception", "proxy_rate_limit", "guardrail_raised"],
|
||||
)
|
||||
def test_attributed_proxy_reject_leaves_deployment_healthy(rejection: Exception):
|
||||
logger = PrometheusLogger()
|
||||
|
||||
logger.set_llm_deployment_failure_metrics(_attributed_failure_kwargs(rejection, rejected_before_routing=True))
|
||||
|
||||
assert logger.litellm_deployment_state._metrics == {}
|
||||
assert _model_id_values(logger.litellm_deployment_failure_responses) == {""}
|
||||
assert _model_id_values(logger.litellm_deployment_total_requests) == {""}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"failure",
|
||||
[
|
||||
_ProviderError("upstream 500"),
|
||||
GuardrailRaisedException(guardrail_name="pii", message="response blocked", status_code=400),
|
||||
],
|
||||
ids=["provider_error", "post_call_guardrail"],
|
||||
)
|
||||
def test_failure_after_routing_still_marks_deployment_partial_outage(failure: Exception):
|
||||
logger = PrometheusLogger()
|
||||
|
||||
logger.set_llm_deployment_failure_metrics(_attributed_failure_kwargs(failure, rejected_before_routing=False))
|
||||
|
||||
assert _model_id_values(logger.litellm_deployment_state) == {"dep-1"}
|
||||
assert _model_id_values(logger.litellm_deployment_failure_responses) == {"dep-1"}
|
||||
|
|
@ -116,6 +116,22 @@ async def test_unknown_models_collapse_to_one_series_on_proxy_request_metrics(ro
|
|||
assert _total_value(metric) == 25
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", [["gpt-4o-mini"], {"name": "gpt-4o-mini"}, 123])
|
||||
async def test_non_string_models_collapse_to_other_on_proxy_request_metrics(router, model: object):
|
||||
logger = PrometheusLogger()
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", router, create=True): # test-quality-ok: production reads proxy_server.llm_router lazily, no injection seam
|
||||
await logger.async_post_call_failure_hook(
|
||||
request_data={"model": model, "metadata": {}, "proxy_server_request": {}},
|
||||
original_exception=_ClientSideError("'model' must be a string."),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key-1"),
|
||||
)
|
||||
|
||||
assert _requested_model_values(logger.litellm_proxy_failed_requests_metric) == {UNRECOGNIZED_REQUESTED_MODEL_LABEL}
|
||||
assert _total_value(logger.litellm_proxy_failed_requests_metric) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_known_alias_and_wildcard_models_keep_their_own_labels(router):
|
||||
logger = PrometheusLogger()
|
||||
|
|
|
|||
|
|
@ -299,42 +299,6 @@ def test_reasoning_tokens_gemini(_local_model_cost_map):
|
|||
)
|
||||
|
||||
|
||||
def test_reasoning_tokens_gemini_3_1_flash_lite(_local_model_cost_map):
|
||||
"""Test cost calculation for gemini-3.1-flash-lite-preview with reasoning tokens"""
|
||||
model = "gemini-3.1-flash-lite-preview"
|
||||
custom_llm_provider = "gemini"
|
||||
|
||||
usage = Usage(
|
||||
completion_tokens=1000,
|
||||
prompt_tokens=500,
|
||||
total_tokens=1500,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=None,
|
||||
audio_tokens=None,
|
||||
reasoning_tokens=400,
|
||||
rejected_prediction_tokens=None,
|
||||
text_tokens=600,
|
||||
),
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
audio_tokens=None, cached_tokens=None, text_tokens=500, image_tokens=None
|
||||
),
|
||||
)
|
||||
model_cost_map = litellm.model_cost[model]
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
assert round(prompt_cost, 10) == round(
|
||||
model_cost_map["input_cost_per_token"] * usage.prompt_tokens,
|
||||
10,
|
||||
)
|
||||
assert round(completion_cost, 10) == round(
|
||||
(model_cost_map["output_cost_per_token"] * usage.completion_tokens_details.text_tokens)
|
||||
+ (model_cost_map["output_cost_per_reasoning_token"] * usage.completion_tokens_details.reasoning_tokens),
|
||||
10,
|
||||
)
|
||||
|
||||
|
||||
def test_image_tokens_with_custom_pricing():
|
||||
|
|
@ -2221,65 +2185,8 @@ def test_vertex_image_generation_cost_falls_back_to_flat_image_pricing(_local_mo
|
|||
assert round(cost, 10) == round(expected_cost, 10)
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_prefers_token_usage_metadata(_local_model_cost_map):
|
||||
"""
|
||||
When usage metadata exists on image responses, Gemini image generation cost
|
||||
should be calculated from token pricing, not flat output_cost_per_image.
|
||||
"""
|
||||
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
||||
input_text_tokens = 20
|
||||
input_image_tokens = 1120
|
||||
output_image_tokens = 1120
|
||||
prompt_tokens = input_text_tokens + input_image_tokens
|
||||
|
||||
image_response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")],
|
||||
usage=ImageUsage(
|
||||
input_tokens=prompt_tokens,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(
|
||||
text_tokens=input_text_tokens,
|
||||
image_tokens=input_image_tokens,
|
||||
),
|
||||
output_tokens=output_image_tokens,
|
||||
total_tokens=prompt_tokens + output_image_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
cost = gemini_image_generation_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
expected_prompt_cost = prompt_tokens * model_info["input_cost_per_token"]
|
||||
expected_completion_cost = output_image_tokens * model_info["output_cost_per_image_token"]
|
||||
expected_total_cost = expected_prompt_cost + expected_completion_cost
|
||||
|
||||
assert round(cost, 10) == round(expected_total_cost, 10)
|
||||
# Ensure this is not falling back to flat per-image pricing.
|
||||
assert cost != len(image_response.data) * model_info["output_cost_per_image"]
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_falls_back_to_flat_image_pricing(_local_model_cost_map):
|
||||
"""
|
||||
Without usage metadata, Gemini image generation cost should fall back to
|
||||
output_cost_per_image * number_of_images.
|
||||
"""
|
||||
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
||||
image_response = ImageResponse(data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")])
|
||||
|
||||
cost = gemini_image_generation_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
expected_cost = len(image_response.data) * model_info["output_cost_per_image"]
|
||||
assert round(cost, 10) == round(expected_cost, 10)
|
||||
|
||||
|
||||
def test_query_count_is_free_without_a_per_query_price(_local_model_cost_map):
|
||||
|
|
@ -2460,23 +2367,6 @@ def test_vertex_global_or_absent_location_no_uplift(vertex_location, _local_mode
|
|||
assert base == located
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["claude-opus-4-1", "gemini-2.0-flash-001"])
|
||||
def test_vertex_location_no_uplift_for_uniformly_priced_model(model, _local_model_cost_map):
|
||||
"""Models Google prices uniformly across endpoints (Gemini 2.x, Claude Opus 4.1
|
||||
and older) carry no multiplier and must not move with the location."""
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
|
||||
|
||||
base = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="vertex_ai")
|
||||
regional = generic_cost_per_token(
|
||||
model=model,
|
||||
usage=usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_location="us-east5",
|
||||
)
|
||||
|
||||
assert base == regional, f"{model} should not have a regional-endpoint uplift"
|
||||
|
||||
|
||||
def test_vertex_uplift_invalid_multiplier_defaults_to_one():
|
||||
|
|
@ -3695,47 +3585,6 @@ def test_route_image_generation_cost_openai_honors_deployment_input_cost_per_ima
|
|||
assert cost == pytest.approx(0.07)
|
||||
|
||||
|
||||
def test_route_image_generation_cost_gemini_adds_grounding_to_deployment_image_price(
|
||||
_local_model_cost_map: None,
|
||||
) -> None:
|
||||
usage = ImageUsage(
|
||||
input_tokens=0,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=0, text_tokens=0),
|
||||
output_tokens=0,
|
||||
total_tokens=0,
|
||||
web_search_requests=3,
|
||||
)
|
||||
|
||||
cost = CostCalculatorUtils.route_image_generation_cost_calculator(
|
||||
model="gemini/gemini-3.1-flash-image-preview",
|
||||
completion_response=_image_response(usage=usage),
|
||||
custom_llm_provider="gemini",
|
||||
call_type="image_generation",
|
||||
model_info={"output_cost_per_image": 0.1},
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(0.1 + 3 * 0.014)
|
||||
|
||||
|
||||
def test_route_image_generation_cost_gemini_bills_tokens_when_no_image_returned(
|
||||
_local_model_cost_map: None,
|
||||
) -> None:
|
||||
usage = ImageUsage(
|
||||
input_tokens=10,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=0, text_tokens=10),
|
||||
output_tokens=1290,
|
||||
total_tokens=1300,
|
||||
)
|
||||
|
||||
cost = CostCalculatorUtils.route_image_generation_cost_calculator(
|
||||
model="gemini/gemini-3.1-flash-image-preview",
|
||||
completion_response=ImageResponse(data=[], usage=usage),
|
||||
custom_llm_provider="gemini",
|
||||
call_type="image_generation",
|
||||
model_info={"output_cost_per_image": 0.08},
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(10 * 5e-07 + 1290 * 6e-05)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
|
|
@ -109,26 +109,6 @@ def test_get_cost_for_built_in_tools_file_search():
|
|||
assert cost == 0.00
|
||||
|
||||
|
||||
def test_get_cost_for_anthropic_web_search():
|
||||
"""
|
||||
Test that Anthropic web search cost is tracked when usage.server_tool_use.web_search_requests
|
||||
is set. Use claude-3-7-sonnet-20250219 (has search_context_cost_per_query) and
|
||||
custom_llm_provider=anthropic so get_cost_for_anthropic_web_search is invoked.
|
||||
"""
|
||||
from litellm.types.utils import ServerToolUse, Usage
|
||||
|
||||
model = "claude-3-7-sonnet-20250219"
|
||||
usage = Usage(server_tool_use=ServerToolUse(web_search_requests=1))
|
||||
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
|
||||
model=model,
|
||||
usage=usage,
|
||||
response_object=None,
|
||||
standard_built_in_tools_params=None,
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
assert cost > 0.0
|
||||
|
||||
|
||||
def test_get_cost_for_anthropic_web_search_with_server_tool_use_dict():
|
||||
"""
|
||||
Anthropic-compatible passthrough responses can construct Usage from a raw
|
||||
|
|
@ -145,88 +125,6 @@ def test_get_cost_for_anthropic_web_search_with_server_tool_use_dict():
|
|||
)
|
||||
|
||||
|
||||
def test_anthropic_web_search_cost_from_raw_response_dict_when_usage_drops_server_tool_use():
|
||||
"""
|
||||
Regression: on the Anthropic /v1/messages sync cost path the response is the raw
|
||||
Anthropic dict while the reconstructed OpenAI-shape Usage drops server_tool_use.
|
||||
The web-search fee must still be charged by reading the count off the raw dict,
|
||||
and the passed-in Usage must not be mutated.
|
||||
"""
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
model = "claude-3-7-sonnet-20250219"
|
||||
web_search_requests = 3
|
||||
raw_response = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": model,
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"server_tool_use": {"web_search_requests": web_search_requests},
|
||||
},
|
||||
}
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
assert getattr(usage, "server_tool_use", None) is None
|
||||
|
||||
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
|
||||
model=model,
|
||||
usage=usage,
|
||||
response_object=raw_response,
|
||||
custom_llm_provider="anthropic",
|
||||
standard_built_in_tools_params=None,
|
||||
)
|
||||
|
||||
per_query_cost = litellm.get_model_info(model)["search_context_cost_per_query"][
|
||||
"search_context_size_medium"
|
||||
]
|
||||
assert cost == per_query_cost * web_search_requests
|
||||
assert cost > 0.0
|
||||
assert getattr(usage, "server_tool_use", None) is None
|
||||
|
||||
|
||||
def test_anthropic_web_search_cost_from_raw_response_dict_when_usage_is_none():
|
||||
"""
|
||||
Regression: when a caller hands the cost tracker a raw Anthropic dict without a
|
||||
parallel Usage object, the web-search fee must still be priced per request from
|
||||
usage.server_tool_use.web_search_requests on the dict instead of falling back to
|
||||
the flat search_context_size_medium tier.
|
||||
"""
|
||||
model = "claude-3-7-sonnet-20250219"
|
||||
web_search_requests = 4
|
||||
raw_response = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": model,
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"server_tool_use": {"web_search_requests": web_search_requests},
|
||||
},
|
||||
}
|
||||
|
||||
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
|
||||
model=model,
|
||||
usage=None,
|
||||
response_object=raw_response,
|
||||
custom_llm_provider="anthropic",
|
||||
standard_built_in_tools_params=None,
|
||||
)
|
||||
|
||||
per_query_cost = litellm.get_model_info(model)["search_context_cost_per_query"][
|
||||
"search_context_size_medium"
|
||||
]
|
||||
assert cost == per_query_cost * web_search_requests
|
||||
|
||||
|
||||
def test_anthropic_web_search_zero_requests_from_raw_response_charges_zero():
|
||||
"""
|
||||
Regression: a raw Anthropic dict reporting zero web search requests must price
|
||||
|
|
@ -287,27 +185,6 @@ def test_anthropic_response_usage_block_preserves_server_tool_use():
|
|||
assert dumped_usage["server_tool_use"] == {"web_search_requests": 2}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model", ["gemini/gemini-2.0-flash-001", "gemini-2.0-flash-001"]
|
||||
)
|
||||
def test_get_cost_for_gemini_web_search(model):
|
||||
"""
|
||||
Test that the cost for a web search is 0.00 when no response object is provided
|
||||
"""
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(web_search_requests=1)
|
||||
)
|
||||
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
|
||||
model=model,
|
||||
usage=usage,
|
||||
response_object=None,
|
||||
standard_built_in_tools_params=None,
|
||||
)
|
||||
assert cost > 0.0
|
||||
|
||||
|
||||
def test_completion_cost_includes_web_search_without_standard_built_in_tools_params():
|
||||
"""
|
||||
Test that completion_cost includes web search cost even when
|
||||
|
|
|
|||
|
|
@ -979,10 +979,6 @@ def test_shipped_tool_search_rule_fills_mapped_claude_entries_without_flag(shipp
|
|||
assert "supports_tool_search" not in litellm.model_cost[key]
|
||||
assert litellm.get_model_info(model, custom_llm_provider=provider)["supports_tool_search"] is True
|
||||
|
||||
assert "supports_tool_search" not in litellm.model_cost["claude-opus-4-1"]
|
||||
opus_4_1_info = litellm.get_model_info("claude-opus-4-1", custom_llm_provider="anthropic")
|
||||
assert opus_4_1_info.get("supports_tool_search") is None
|
||||
|
||||
assert "supports_tool_search" not in litellm.model_cost["azure_ai/claude-opus-5"]
|
||||
azure_opus_5_info = litellm.get_model_info("claude-opus-5", custom_llm_provider="azure_ai")
|
||||
assert azure_opus_5_info.get("supports_tool_search") is None
|
||||
|
|
|
|||
|
|
@ -218,7 +218,6 @@ def test_shipped_backup_marks_claude_4_6_plus_adaptive_not_4_0():
|
|||
assert backup[adaptive]["supports_adaptive_thinking"] is True, adaptive
|
||||
|
||||
for non_adaptive in [
|
||||
"claude-opus-4-20250514",
|
||||
"us.anthropic.claude-opus-4-20250514-v1:0",
|
||||
"claude-opus-4-5",
|
||||
]:
|
||||
|
|
|
|||
|
|
@ -427,7 +427,7 @@ async def test_async_inline_remote_media_cancels_the_other_fetches_when_one_fail
|
|||
_SSRF_VERDICTS = (
|
||||
SSRFError(
|
||||
"URL targets a blocked address (10.0.0.8). If this is a legitimate internal service, "
|
||||
"add the host to `user_url_allowed_hosts` in general_settings."
|
||||
"add the host to `user_url_allowed_hosts` in litellm_settings."
|
||||
),
|
||||
SSRFError("DNS resolution failed for 'internal.example': [Errno 8] nodename nor servname provided, or not known"),
|
||||
SSRFError("No addresses found for 'internal.example'"),
|
||||
|
|
|
|||
|
|
@ -8389,21 +8389,6 @@ def test_get_assembled_streaming_response_bills_a_provider_reported_usage_cost()
|
|||
assert logging_obj._response_cost_calculator(result=assembled) == 0.0042
|
||||
|
||||
|
||||
def test_get_assembled_streaming_response_without_usage_cost_leaves_pricing_to_the_price_map():
|
||||
logging_obj = _responses_stream_logging_obj()
|
||||
now = datetime.datetime.now()
|
||||
|
||||
assembled = logging_obj._get_assembled_streaming_response(
|
||||
result=_completed_responses_event(ResponseAPIUsage(input_tokens=12, output_tokens=2, total_tokens=14)),
|
||||
start_time=now,
|
||||
end_time=now,
|
||||
is_async=True,
|
||||
streaming_chunks=[],
|
||||
)
|
||||
|
||||
assert "additional_headers" not in assembled._hidden_params
|
||||
price_map_cost = logging_obj._response_cost_calculator(result=assembled)
|
||||
assert price_map_cost is not None and 0 < price_map_cost != 0.0042
|
||||
|
||||
|
||||
def test_response_cost_calculator_prices_terminal_responses_event_from_its_response():
|
||||
|
|
|
|||
|
|
@ -99,29 +99,3 @@ def test_stream_chunk_builder_coerces_server_tool_use_to_pydantic():
|
|||
assert server_tool_use.web_search_requests == 3
|
||||
|
||||
|
||||
def test_completion_cost_does_not_raise_on_streaming_web_search_response():
|
||||
"""
|
||||
Regression: completion_cost(...) must not raise AttributeError when the
|
||||
response was reconstructed by stream_chunk_builder from a streaming
|
||||
Anthropic web_search call.
|
||||
"""
|
||||
chunks = [
|
||||
_make_text_chunk("hello"),
|
||||
_make_finish_chunk_with_usage_dict_server_tool_use(),
|
||||
]
|
||||
|
||||
rebuilt = stream_chunk_builder(chunks)
|
||||
assert rebuilt is not None
|
||||
|
||||
# The exact dollar amount depends on the model-pricing table; what matters
|
||||
# for this regression is that it does NOT raise AttributeError on
|
||||
# `dict has no attribute 'web_search_requests'`.
|
||||
try:
|
||||
cost = completion_cost(completion_response=rebuilt)
|
||||
except AttributeError as e: # pragma: no cover - regression guard
|
||||
pytest.fail(
|
||||
"completion_cost raised AttributeError after stream_chunk_builder "
|
||||
f"(issue #26153 regression): {e}"
|
||||
)
|
||||
|
||||
assert isinstance(cost, (int, float))
|
||||
|
|
|
|||
|
|
@ -633,9 +633,6 @@ def test_openai_token_with_image_and_text():
|
|||
"model, base_model, input_tokens, user_max_tokens, expected_value",
|
||||
[
|
||||
("random-model", "random-model", 1024, 1024, 1024),
|
||||
("command", "command", 1000000, None, None), # model max = 4096
|
||||
("command", "command", 4000, 256, 96), # model max = 4096
|
||||
("command", "command", 4000, 10, 10), # model max = 4096
|
||||
("gpt-3.5-turbo", "gpt-3.5-turbo", 4000, 5000, 4096), # model max output = 4096
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -5580,43 +5580,6 @@ def test_tool_config_cachepoint_not_placed_or_credited_for_model_without_prompt_
|
|||
assert "litellm_gateway_injected_cache" not in bucket
|
||||
|
||||
|
||||
def test_translate_response_format_json_schema_still_injects_tool():
|
||||
"""
|
||||
response_format with an explicit json_schema should still use the
|
||||
synthetic tool call approach (for models that don't support native
|
||||
structured outputs).
|
||||
"""
|
||||
config = AmazonConverseConfig()
|
||||
|
||||
response_format = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "FactResult",
|
||||
"schema": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"facts": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"},
|
||||
},
|
||||
},
|
||||
"required": ["facts"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
optional_params: dict = {}
|
||||
result = config._translate_response_format_param(
|
||||
value=response_format,
|
||||
model="anthropic.claude-3-haiku-20240307-v1:0",
|
||||
optional_params=optional_params,
|
||||
non_default_params={"response_format": response_format},
|
||||
is_thinking_enabled=False,
|
||||
)
|
||||
|
||||
assert result["json_mode"] is True
|
||||
assert "tools" in result
|
||||
assert "tool_choice" in result
|
||||
|
||||
|
||||
def test_transform_response_finish_reason_stop_when_json_mode_filters_all_tools():
|
||||
|
|
|
|||
|
|
@ -153,14 +153,6 @@ def test_uncached_request_bills_every_prompt_token_at_the_input_rate(local_model
|
|||
assert completion_cost == pytest.approx(200 * info["output_cost_per_token"])
|
||||
|
||||
|
||||
def test_legacy_endpoint_names_still_resolve(local_model_cost_map: None) -> None:
|
||||
info: Final = _model_info("databricks/databricks-mixtral-8x7b-instruct")
|
||||
usage: Final = Usage(prompt_tokens=100, completion_tokens=100, total_tokens=200)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="databricks/mixtral-8x7b-instruct-v0.1", usage=usage)
|
||||
|
||||
assert prompt_cost == pytest.approx(100 * info["input_cost_per_token"])
|
||||
assert completion_cost == pytest.approx(100 * info["output_cost_per_token"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", NEW_MODELS)
|
||||
|
|
|
|||
|
|
@ -200,172 +200,12 @@ def test_maps_no_usage_details():
|
|||
assert cost_per_google_maps_grounding_request(usage=usage, model_info=model_info) == 0.0
|
||||
|
||||
|
||||
def test_gemini_image_edit_cost_prefers_token_usage_metadata(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
||||
input_text_tokens = 20
|
||||
input_image_tokens = 1120
|
||||
output_image_tokens = 1120
|
||||
prompt_tokens = input_text_tokens + input_image_tokens
|
||||
image_response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")],
|
||||
usage=ImageUsage(
|
||||
input_tokens=prompt_tokens,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(
|
||||
text_tokens=input_text_tokens,
|
||||
image_tokens=input_image_tokens,
|
||||
),
|
||||
output_tokens=output_image_tokens,
|
||||
total_tokens=prompt_tokens + output_image_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
cost = gemini_image_edit_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
expected_cost = (
|
||||
prompt_tokens * model_info["input_cost_per_token"]
|
||||
+ output_image_tokens * model_info["output_cost_per_image_token"]
|
||||
)
|
||||
flat_image_cost = (
|
||||
len(image_response.data or []) * model_info["output_cost_per_image"]
|
||||
)
|
||||
assert round(cost, 10) == round(expected_cost, 10)
|
||||
assert cost != flat_image_cost
|
||||
|
||||
|
||||
def test_gemini_image_edit_cost_uses_output_token_details(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
||||
input_text_tokens = 20
|
||||
output_text_tokens = 213
|
||||
output_image_tokens = 1120
|
||||
output_tokens = output_text_tokens + output_image_tokens
|
||||
image_response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img1")],
|
||||
usage=ImageUsage(
|
||||
input_tokens=input_text_tokens,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(
|
||||
text_tokens=input_text_tokens,
|
||||
image_tokens=0,
|
||||
),
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=input_text_tokens + output_tokens,
|
||||
prompt_tokens=input_text_tokens,
|
||||
completion_tokens=output_tokens,
|
||||
prompt_tokens_details={
|
||||
"text_tokens": input_text_tokens,
|
||||
"image_tokens": 0,
|
||||
},
|
||||
completion_tokens_details={
|
||||
"text_tokens": output_text_tokens,
|
||||
"image_tokens": output_image_tokens,
|
||||
},
|
||||
output_tokens_details={
|
||||
"text_tokens": output_text_tokens,
|
||||
"image_tokens": output_image_tokens,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
cost = gemini_image_edit_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
expected_cost = (
|
||||
input_text_tokens * model_info["input_cost_per_token"]
|
||||
+ output_text_tokens * model_info["output_cost_per_token"]
|
||||
+ output_image_tokens * model_info["output_cost_per_image_token"]
|
||||
)
|
||||
all_output_as_image_cost = (
|
||||
input_text_tokens * model_info["input_cost_per_token"]
|
||||
+ (output_text_tokens + output_image_tokens)
|
||||
* model_info["output_cost_per_image_token"]
|
||||
)
|
||||
assert round(cost, 10) == round(expected_cost, 10)
|
||||
assert cost != all_output_as_image_cost
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_uses_output_token_details(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
||||
input_text_tokens = 20
|
||||
output_text_tokens = 213
|
||||
output_image_tokens = 1120
|
||||
output_tokens = output_text_tokens + output_image_tokens
|
||||
image_response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img1")],
|
||||
usage=ImageUsage(
|
||||
input_tokens=input_text_tokens,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(
|
||||
text_tokens=input_text_tokens,
|
||||
image_tokens=0,
|
||||
),
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=input_text_tokens + output_tokens,
|
||||
prompt_tokens=input_text_tokens,
|
||||
completion_tokens=output_tokens,
|
||||
prompt_tokens_details={
|
||||
"text_tokens": input_text_tokens,
|
||||
"image_tokens": 0,
|
||||
},
|
||||
completion_tokens_details={
|
||||
"text_tokens": output_text_tokens,
|
||||
"image_tokens": output_image_tokens,
|
||||
},
|
||||
output_tokens_details={
|
||||
"text_tokens": output_text_tokens,
|
||||
"image_tokens": output_image_tokens,
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
cost = gemini_image_generation_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
expected_cost = (
|
||||
input_text_tokens * model_info["input_cost_per_token"]
|
||||
+ output_text_tokens * model_info["output_cost_per_token"]
|
||||
+ output_image_tokens * model_info["output_cost_per_image_token"]
|
||||
)
|
||||
all_output_as_image_cost = (
|
||||
input_text_tokens * model_info["input_cost_per_token"]
|
||||
+ (output_text_tokens + output_image_tokens)
|
||||
* model_info["output_cost_per_image_token"]
|
||||
)
|
||||
assert round(cost, 10) == round(expected_cost, 10)
|
||||
assert cost != all_output_as_image_cost
|
||||
|
||||
|
||||
def test_gemini_image_edit_cost_falls_back_to_flat_image_pricing(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
image_response = ImageResponse(
|
||||
data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")]
|
||||
)
|
||||
|
||||
cost = gemini_image_edit_cost_calculator(
|
||||
model=model,
|
||||
image_response=image_response,
|
||||
)
|
||||
|
||||
assert cost == len(image_response.data or []) * model_info["output_cost_per_image"]
|
||||
|
||||
|
||||
def _image_response_with_web_search(web_search_requests):
|
||||
|
|
@ -383,43 +223,8 @@ def _image_response_with_web_search(web_search_requests):
|
|||
return ImageResponse(data=[ImageObject(b64_json="img1")], usage=usage)
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_adds_web_search_grounding(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
||||
grounded = gemini_image_generation_cost_calculator(
|
||||
model=model,
|
||||
image_response=_image_response_with_web_search(2),
|
||||
)
|
||||
ungrounded = gemini_image_generation_cost_calculator(
|
||||
model=model,
|
||||
image_response=_image_response_with_web_search(None),
|
||||
)
|
||||
|
||||
expected_web_search_cost = cost_per_web_search_request(
|
||||
usage=_make_usage(2), model_info=model_info
|
||||
)
|
||||
assert expected_web_search_cost > 0
|
||||
assert round(grounded - ungrounded, 10) == round(expected_web_search_cost, 10)
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_no_web_search_when_absent(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
|
||||
cost_zero = gemini_image_generation_cost_calculator(
|
||||
model=model,
|
||||
image_response=_image_response_with_web_search(0),
|
||||
)
|
||||
cost_none = gemini_image_generation_cost_calculator(
|
||||
model=model,
|
||||
image_response=_image_response_with_web_search(None),
|
||||
)
|
||||
|
||||
assert cost_zero == cost_none
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
|
|
@ -371,62 +371,8 @@ def test_x_initiator_header_system_only_messages():
|
|||
assert headers["X-Initiator"] == "user"
|
||||
|
||||
|
||||
def test_get_supported_openai_params_claude_model():
|
||||
"""Test that Claude models with extended thinking support have thinking and reasoning parameters."""
|
||||
config = GithubCopilotConfig()
|
||||
|
||||
# Test Claude 4 model supports thinking and reasoning_effort parameters
|
||||
supported_params = config.get_supported_openai_params("claude-sonnet-4-20250514")
|
||||
assert "thinking" in supported_params
|
||||
assert "reasoning_effort" in supported_params
|
||||
|
||||
# Test Claude 3-7 model supports thinking and reasoning_effort parameters
|
||||
supported_params_claude37 = config.get_supported_openai_params(
|
||||
"claude-3-7-sonnet-20250219"
|
||||
)
|
||||
assert "thinking" in supported_params_claude37
|
||||
assert "reasoning_effort" in supported_params_claude37
|
||||
|
||||
# Test Claude 3.5 model does NOT support thinking parameters (no extended thinking)
|
||||
supported_params_claude35 = config.get_supported_openai_params("claude-3.5-sonnet")
|
||||
assert "thinking" not in supported_params_claude35
|
||||
assert "reasoning_effort" not in supported_params_claude35
|
||||
|
||||
# Test non-Claude model doesn't include thinking parameters but may include reasoning_effort
|
||||
supported_params_gpt = config.get_supported_openai_params("gpt-4o")
|
||||
assert "thinking" not in supported_params_gpt
|
||||
# gpt-4o should NOT have reasoning_effort (not a reasoning model)
|
||||
assert "reasoning_effort" not in supported_params_gpt
|
||||
|
||||
# Test O-series reasoning models include reasoning_effort but not thinking
|
||||
supported_params_o3 = config.get_supported_openai_params("o3-mini")
|
||||
assert "thinking" not in supported_params_o3
|
||||
# o3-mini should have reasoning_effort (it's an O-series reasoning model)
|
||||
assert "reasoning_effort" in supported_params_o3
|
||||
|
||||
|
||||
def test_get_supported_openai_params_case_insensitive():
|
||||
"""Test that Claude model detection is case-insensitive for models with extended thinking."""
|
||||
config = GithubCopilotConfig()
|
||||
|
||||
# Test uppercase Claude 4 model with full model name
|
||||
supported_params_upper = config.get_supported_openai_params(
|
||||
"CLAUDE-SONNET-4-20250514"
|
||||
)
|
||||
assert "thinking" in supported_params_upper
|
||||
assert "reasoning_effort" in supported_params_upper
|
||||
|
||||
# Test mixed case Claude 3-7 model (has extended thinking) with full model name
|
||||
supported_params_mixed = config.get_supported_openai_params(
|
||||
"Claude-3-7-Sonnet-20250219"
|
||||
)
|
||||
assert "thinking" in supported_params_mixed
|
||||
assert "reasoning_effort" in supported_params_mixed
|
||||
|
||||
# Test that Claude 3.5 models don't have thinking support (case insensitive)
|
||||
supported_params_35 = config.get_supported_openai_params("CLAUDE-3.5-SONNET")
|
||||
assert "thinking" not in supported_params_35
|
||||
assert "reasoning_effort" not in supported_params_35
|
||||
|
||||
|
||||
def test_copilot_vision_request_header_with_image():
|
||||
|
|
|
|||
|
|
@ -38,10 +38,6 @@ def test_gpt5_supports_reasoning_effort(config: OpenAIConfig):
|
|||
assert "reasoning_effort" in config.get_supported_openai_params(model="gpt-5-mini")
|
||||
|
||||
|
||||
def test_gpt5_chat_does_not_support_reasoning_effort(config: OpenAIConfig):
|
||||
assert "reasoning_effort" not in config.get_supported_openai_params(
|
||||
model="gpt-5-chat-latest"
|
||||
)
|
||||
|
||||
|
||||
def test_gpt5_chat_supports_temperature(config: OpenAIConfig):
|
||||
|
|
@ -174,10 +170,6 @@ def test_gpt5_codex_unsupported_params_drop(config: OpenAIConfig):
|
|||
assert param not in config.get_supported_openai_params(model="gpt-5-codex")
|
||||
|
||||
|
||||
def test_gpt5_codex_supports_tool_choice(gpt5_config: OpenAIGPT5Config):
|
||||
"""Test that GPT-5-Codex supports tool_choice parameter."""
|
||||
supported_params = gpt5_config.get_supported_openai_params(model="gpt-5-codex")
|
||||
assert "tool_choice" in supported_params
|
||||
|
||||
|
||||
def test_gpt5_codex_supports_function_calling(config: OpenAIConfig):
|
||||
|
|
@ -246,14 +238,6 @@ def test_gpt5_1_reasoning_effort_none(config: OpenAIConfig):
|
|||
assert params["reasoning_effort"] == effort
|
||||
|
||||
|
||||
def test_gpt5_1_codex_max_allows_reasoning_effort_xhigh(config: OpenAIConfig):
|
||||
params = config.map_openai_params(
|
||||
non_default_params={"reasoning_effort": "xhigh"},
|
||||
optional_params={},
|
||||
model="gpt-5.1-codex-max",
|
||||
drop_params=False,
|
||||
)
|
||||
assert params["reasoning_effort"] == "xhigh"
|
||||
|
||||
|
||||
def test_gpt5_rejects_reasoning_effort_xhigh_for_other_models(config: OpenAIConfig):
|
||||
|
|
|
|||
|
|
@ -5948,8 +5948,8 @@ def test_calculate_web_search_requests_counts_unique_queries():
|
|||
@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai"])
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
["gemini-2.5-flash", "gemini-3-pro-preview"],
|
||||
ids=["thinking_budget_mapper", "thinking_level_mapper"],
|
||||
["gemini-2.5-flash"],
|
||||
ids=["thinking_budget_mapper"],
|
||||
)
|
||||
@pytest.mark.parametrize("reasoning_effort", ["banana", "xhigh"])
|
||||
def test_invalid_reasoning_effort_is_a_400_not_a_500(custom_llm_provider, model, reasoning_effort):
|
||||
|
|
|
|||
|
|
@ -238,31 +238,3 @@ def test_audio_predict_response_supports_bytes_base64_encoded(
|
|||
assert logging_obj.model_call_details["response_cost"] == pytest.approx(0.06)
|
||||
|
||||
|
||||
def test_image_predict_response_is_not_billed_as_audio(
|
||||
local_model_cost_map: None,
|
||||
) -> None:
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json={"predictions": [{"bytesBase64Encoded": "frame", "mimeType": "image/png"}]},
|
||||
)
|
||||
|
||||
result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
|
||||
httpx_response=response,
|
||||
logging_obj=logging_obj,
|
||||
url_route=(
|
||||
"/v1/projects/test/locations/us-central1/publishers/google/models/imagen-4.0-generate-001:predict"
|
||||
),
|
||||
result=response.text,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={"instances": [{"prompt": "a red cube"}]},
|
||||
)
|
||||
|
||||
assert isinstance(result["result"], litellm.ImageResponse)
|
||||
assert logging_obj.call_type == PassthroughCallTypes.passthrough_image_generation.value
|
||||
assert result["kwargs"]["response_cost"] == pytest.approx(
|
||||
litellm.model_cost["vertex_ai/imagen-4.0-generate-001"]["output_cost_per_image"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -37,10 +37,6 @@ WANDB_REASONING_MODELS: Final = (
|
|||
"Qwen/Qwen3.5-35B-A3B",
|
||||
"zai-org/GLM-5.2",
|
||||
"moonshotai/Kimi-K2.5",
|
||||
"MiniMaxAI/MiniMax-M2.5",
|
||||
"zai-org/GLM-4.5",
|
||||
"Qwen/Qwen3-235B-A22B-Thinking-2507",
|
||||
"deepseek-ai/DeepSeek-R1-0528",
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -152,83 +152,10 @@ class TestXAICostCalculator:
|
|||
setattr(reported, "server_side_tool_usage_details", {"web_search_calls": 3})
|
||||
assert get_cost_for_web_search_request("xai", reported, {}) == 0.0
|
||||
|
||||
def test_no_reported_cost_falls_back_to_token_math(self):
|
||||
"""Absent the provider figure, nothing changes for existing callers."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
|
||||
|
||||
assert prompt_cost > 0.0
|
||||
assert completion_cost > 0.0
|
||||
|
||||
def test_malformed_reported_cost_falls_back_to_token_math(self):
|
||||
"""A junk value must not fail the request, fall back to calculating."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
|
||||
setattr(usage, "cost", "not-a-number")
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
|
||||
|
||||
assert prompt_cost > 0.0
|
||||
assert completion_cost > 0.0
|
||||
|
||||
def test_boolean_reported_cost_falls_back_to_token_math(self):
|
||||
"""True is an int in python and would otherwise be billed as $1."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
|
||||
setattr(usage, "cost", True)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
|
||||
|
||||
assert prompt_cost > 0.0
|
||||
assert completion_cost > 0.0
|
||||
assert completion_cost != 1.0
|
||||
|
||||
def test_negative_reported_cost_is_rejected(self):
|
||||
"""A negative amount must never reach spend tracking.
|
||||
|
||||
A caller who can set api_base controls the response body, so trusting a
|
||||
negative figure would let them subtract from their own recorded spend and
|
||||
slip past a budget. Fall back to token pricing instead, and keep charging
|
||||
the web search surcharge, since no trustworthy total was reported.
|
||||
"""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=200,
|
||||
total_tokens=300,
|
||||
cost=-0.0037756,
|
||||
)
|
||||
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 3})
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
|
||||
|
||||
assert prompt_cost > 0.0
|
||||
assert completion_cost > 0.0
|
||||
assert cost_per_web_search_request(usage=usage, model_info={}) > 0.0
|
||||
|
||||
def test_non_finite_reported_cost_is_rejected(self):
|
||||
"""NaN compares false against every budget threshold.
|
||||
|
||||
Usage stores a provider supplied cost without validating it, so a caller who
|
||||
controls the response body could report NaN and leave spend >= max_budget
|
||||
false for the life of the key rather than mispricing one request. The
|
||||
infinities are refused alongside it. Fall back to token pricing and keep
|
||||
charging the web search surcharge, since no trustworthy total was reported.
|
||||
"""
|
||||
for reported_cost in (float("nan"), float("inf"), float("-inf")):
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=200,
|
||||
total_tokens=300,
|
||||
cost=reported_cost,
|
||||
)
|
||||
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 3})
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
|
||||
|
||||
assert math.isfinite(prompt_cost), reported_cost
|
||||
assert math.isfinite(completion_cost), reported_cost
|
||||
assert prompt_cost > 0.0, reported_cost
|
||||
assert completion_cost > 0.0, reported_cost
|
||||
assert cost_per_web_search_request(usage=usage, model_info={}) > 0.0, reported_cost
|
||||
|
||||
def test_zero_reported_cost_is_honoured(self):
|
||||
"""A reported zero is a real answer, not a missing value."""
|
||||
|
|
|
|||
|
|
@ -1,102 +0,0 @@
|
|||
"""
|
||||
xAI retired eight slugs on 2026-05-15 but kept them resolvable: chat slugs redirect to
|
||||
grok-4.3 and bill at grok-4.3's rates, while the grok-code-fast slugs are aliases of
|
||||
grok-build-0.1 and bill at its rates, so the registry must price them that way or spend
|
||||
tracking is wrong. The grok-3-beta, grok-3-fast, grok-3-mini, and grok-4-1-fast slugs
|
||||
are absent from /v1/language-models and resolve to grok-4.3 the same way (the chat
|
||||
response names grok-4.3 as the served model), so they carry grok-4.3's rates too.
|
||||
https://docs.x.ai/developers/migration/may-15-retirement
|
||||
https://docs.x.ai/developers/models/grok-build-0.1
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
REPO_ROOT = Path(__file__).parents[4]
|
||||
PRICES_PATH = REPO_ROOT / "model_prices_and_context_window.json"
|
||||
BACKUP_PRICES_PATH = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json"
|
||||
MAP_PATHS = (PRICES_PATH, BACKUP_PRICES_PATH)
|
||||
|
||||
REDIRECT_TARGET = "xai/grok-4.3"
|
||||
GROK_3_MINI_SLUGS = (
|
||||
"xai/grok-3-mini",
|
||||
"xai/grok-3-mini-beta",
|
||||
"xai/grok-3-mini-fast",
|
||||
"xai/grok-3-mini-fast-beta",
|
||||
"xai/grok-3-mini-fast-latest",
|
||||
"xai/grok-3-mini-latest",
|
||||
)
|
||||
REDIRECTED_SLUGS = (
|
||||
"xai/grok-3",
|
||||
"xai/grok-3-beta",
|
||||
"xai/grok-3-fast-beta",
|
||||
"xai/grok-3-fast-latest",
|
||||
"xai/grok-3-latest",
|
||||
*GROK_3_MINI_SLUGS,
|
||||
"xai/grok-4",
|
||||
"xai/grok-4-0709",
|
||||
"xai/grok-4-1-fast",
|
||||
"xai/grok-4-1-fast-non-reasoning",
|
||||
"xai/grok-4-1-fast-non-reasoning-latest",
|
||||
"xai/grok-4-1-fast-reasoning",
|
||||
"xai/grok-4-1-fast-reasoning-latest",
|
||||
"xai/grok-4-fast-non-reasoning",
|
||||
"xai/grok-4-fast-reasoning",
|
||||
"xai/grok-4-latest",
|
||||
)
|
||||
CODE_REDIRECT_TARGET = "xai/grok-build-0.1"
|
||||
CODE_SLUGS = (
|
||||
"xai/grok-code-fast",
|
||||
"xai/grok-code-fast-1",
|
||||
"xai/grok-code-fast-1-0825",
|
||||
)
|
||||
BASE_COST_FIELDS = ("input_cost_per_token", "output_cost_per_token", "cache_read_input_token_cost")
|
||||
TIER_COST_FIELDS = (
|
||||
"input_cost_per_token_above_200k_tokens",
|
||||
"output_cost_per_token_above_200k_tokens",
|
||||
"cache_read_input_token_cost_above_200k_tokens",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", params=[p.name for p in MAP_PATHS])
|
||||
def cost_map(request: pytest.FixtureRequest) -> dict:
|
||||
path = next(p for p in MAP_PATHS if p.name == request.param)
|
||||
return json.loads(path.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("slug", REDIRECTED_SLUGS)
|
||||
def test_redirected_slug_bills_at_the_target_rate(cost_map: dict, slug: str):
|
||||
target = cost_map[REDIRECT_TARGET]
|
||||
entry = cost_map[slug]
|
||||
for field in BASE_COST_FIELDS:
|
||||
assert entry[field] == target[field], field
|
||||
|
||||
|
||||
@pytest.mark.parametrize("slug", CODE_SLUGS)
|
||||
def test_code_slug_bills_at_grok_build_rate(cost_map: dict, slug: str):
|
||||
"""grok-code-fast* are aliases of grok-build-0.1, not grok-4.3 redirects."""
|
||||
target = cost_map[CODE_REDIRECT_TARGET]
|
||||
entry = cost_map[slug]
|
||||
for field in (*BASE_COST_FIELDS, *TIER_COST_FIELDS):
|
||||
assert entry[field] == target[field], field
|
||||
|
||||
|
||||
@pytest.mark.parametrize("slug", REDIRECTED_SLUGS)
|
||||
def test_redirected_slug_carries_the_target_tier_rates(cost_map: dict, slug: str):
|
||||
"""The request executes as grok-4.3, so it is tiered at grok-4.3's 200k boundary."""
|
||||
target = cost_map[REDIRECT_TARGET]
|
||||
entry = cost_map[slug]
|
||||
for field in TIER_COST_FIELDS:
|
||||
assert entry[field] == target[field], field
|
||||
assert {k for k in entry if "_above_" in k} == {k for k in target if "_above_" in k}
|
||||
|
||||
|
||||
def test_both_cost_maps_agree_on_the_redirected_slugs():
|
||||
prices = json.loads(PRICES_PATH.read_text(encoding="utf-8"))
|
||||
backup = json.loads(BACKUP_PRICES_PATH.read_text(encoding="utf-8"))
|
||||
for slug in (*REDIRECTED_SLUGS, *CODE_SLUGS, REDIRECT_TARGET, CODE_REDIRECT_TARGET):
|
||||
assert prices[slug] == backup[slug], slug
|
||||
|
|
@ -4,7 +4,7 @@ import pytest
|
|||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
_upstream_credential_headers,
|
||||
upstream_credential_headers,
|
||||
build_synthetic_mcp_request,
|
||||
logging_safe_mcp_headers,
|
||||
validate_and_normalize_mcp_server_payload,
|
||||
|
|
@ -189,8 +189,8 @@ class TestLoggingSafeMcpHeaders:
|
|||
"""clean_headers already strips authorization, and claiming it here would change which
|
||||
header authenticated_with_header resolves to on a config that lists it by design."""
|
||||
with _configured_servers(_server_forwarding("Authorization", "X-GitHub-Token")):
|
||||
assert "authorization" not in _upstream_credential_headers(["authorization", "x-github-token"])
|
||||
assert "x-github-token" in _upstream_credential_headers(["authorization", "x-github-token"])
|
||||
assert "authorization" not in upstream_credential_headers(["authorization", "x-github-token"])
|
||||
assert "x-github-token" in upstream_credential_headers(["authorization", "x-github-token"])
|
||||
|
||||
def test_keeps_headers_when_no_server_forwards_them(self):
|
||||
with _configured_servers(_server_forwarding("x-github-token")):
|
||||
|
|
|
|||
|
|
@ -8057,7 +8057,6 @@ def test_model_has_no_cost_mapping_no_model_or_router_is_false():
|
|||
[
|
||||
"azure/speech/azure-tts",
|
||||
"mistral/mistral-ocr-latest",
|
||||
"vertex_ai/imagen-3.0-generate-001",
|
||||
"dashscope/qwen-flash",
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -122,6 +122,13 @@ class TestPlanPgBouncer:
|
|||
"pgbouncer": "true",
|
||||
}
|
||||
|
||||
def test_an_upstream_that_already_disables_prepared_statements_gets_a_single_pgbouncer_flag(self):
|
||||
pooled: Final = _plan("postgresql://app:pw@db/litellm?connection_limit=5&pgbouncer=true").pooled_url
|
||||
assert urllib.parse.parse_qsl(urllib.parse.urlsplit(pooled).query) == [
|
||||
("connection_limit", "5"),
|
||||
("pgbouncer", "true"),
|
||||
]
|
||||
|
||||
@pytest.mark.parametrize("hop_param", ["channel_binding=require", "gssencmode=require"])
|
||||
def test_transport_params_for_the_postgres_hop_stay_off_the_plain_tcp_loopback_url(self, hop_param: str):
|
||||
pooled: Final = _plan(f"postgresql://app:pw@db/litellm?connection_limit=5&{hop_param}").pooled_url
|
||||
|
|
|
|||
|
|
@ -588,43 +588,6 @@ class TestAzureAnthropicCostCalculation:
|
|||
== "claude-3-5-haiku-20241022"
|
||||
)
|
||||
|
||||
def test_passthrough_logging_sets_response_cost_with_server_tool_use_dict(self):
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
logging_obj = self._create_mock_logging_obj(model="claude-3-7-sonnet-20250219")
|
||||
logging_obj.get_router_model_id.return_value = None
|
||||
logging_obj.litellm_params = {}
|
||||
|
||||
response = ModelResponse(
|
||||
id="test-id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="test", role="assistant"),
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
usage={
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
"server_tool_use": {"web_search_requests": 1},
|
||||
},
|
||||
)
|
||||
|
||||
kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
|
||||
litellm_model_response=response,
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
kwargs={},
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert "response_cost" in kwargs
|
||||
assert kwargs["response_cost"] > 0
|
||||
|
||||
|
||||
class TestAnthropicBatchPassthroughCostTracking:
|
||||
|
|
@ -2355,42 +2318,6 @@ class TestAnthropicResponseCostRecordedOnModelCallDetails:
|
|||
model_call_details["response_cost"], not from kwargs, so the streaming payload
|
||||
builder must record it there or streaming pass-through logs $0."""
|
||||
|
||||
def test_create_payload_records_response_cost_on_model_call_details(self):
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
logging_obj.get_router_model_id.return_value = None
|
||||
logging_obj.litellm_params = {}
|
||||
logging_obj.litellm_call_id = "test-call-id"
|
||||
|
||||
response = ModelResponse(
|
||||
id="test-id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="hello", role="assistant"),
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
usage={"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
||||
)
|
||||
|
||||
kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
|
||||
litellm_model_response=response,
|
||||
model="claude-3-7-sonnet-20250219",
|
||||
kwargs={},
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert (
|
||||
logging_obj.model_call_details["response_cost"] == kwargs["response_cost"]
|
||||
)
|
||||
assert logging_obj.model_call_details["response_cost"] > 0
|
||||
|
||||
|
||||
class TestAnthropicPassthroughFastMode:
|
||||
|
|
|
|||
|
|
@ -462,30 +462,6 @@ class TestVertexAIBatchPassthroughHandler:
|
|||
assert mock_store.call_args[1]["unified_object_id"]
|
||||
assert mock_store.call_args[1]["is_batch_create"] is expected
|
||||
|
||||
def test_batch_cost_calculation_integration(self):
|
||||
"""Single Vertex AI response → non-zero cost with correct token counts."""
|
||||
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
|
||||
|
||||
vertex_ai_batch_responses = [
|
||||
{
|
||||
"response": {
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 10,
|
||||
"candidatesTokenCount": 5,
|
||||
"totalTokenCount": 15,
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
result = calculate_vertex_ai_batch_cost_and_usage(
|
||||
vertex_ai_batch_responses, model_name="gemini-2.0-flash-001"
|
||||
)
|
||||
|
||||
assert result.usage.total_tokens == 15
|
||||
assert result.usage.prompt_tokens == 10
|
||||
assert result.usage.completion_tokens == 5
|
||||
assert result.cost > 0, "batch_cost_calculator should return a non-zero cost"
|
||||
|
||||
def test_batch_response_transformation(self):
|
||||
"""Test transformation of Vertex AI batch responses to OpenAI format"""
|
||||
|
|
@ -639,76 +615,7 @@ class TestVertexAIBatchCostCalculation:
|
|||
batch_cost_calculator — no VertexGeminiConfig transformation involved.
|
||||
"""
|
||||
|
||||
def test_should_aggregate_cost_and_usage_across_responses(self):
|
||||
"""Two successful responses → costs and token counts are summed."""
|
||||
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
|
||||
|
||||
responses = [
|
||||
{
|
||||
"response": {
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 10,
|
||||
"candidatesTokenCount": 5,
|
||||
"totalTokenCount": 15,
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 8,
|
||||
"candidatesTokenCount": 3,
|
||||
"totalTokenCount": 11,
|
||||
}
|
||||
}
|
||||
},
|
||||
]
|
||||
|
||||
result = calculate_vertex_ai_batch_cost_and_usage(
|
||||
responses, model_name="gemini-2.0-flash-001"
|
||||
)
|
||||
|
||||
assert result.usage.prompt_tokens == 18
|
||||
assert result.usage.completion_tokens == 8
|
||||
assert result.usage.total_tokens == 26
|
||||
assert result.cost > 0, "batch_cost_calculator should return a non-zero cost"
|
||||
|
||||
def test_should_skip_responses_with_null_response_body(self):
|
||||
"""Failed lines (response: None) are skipped without error."""
|
||||
from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage
|
||||
|
||||
responses = [
|
||||
{
|
||||
"response": {
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 10,
|
||||
"candidatesTokenCount": 5,
|
||||
"totalTokenCount": 15,
|
||||
}
|
||||
}
|
||||
},
|
||||
{"status": "JOB_STATE_FAILED", "response": None},
|
||||
{
|
||||
"response": {
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 8,
|
||||
"candidatesTokenCount": 3,
|
||||
"totalTokenCount": 11,
|
||||
}
|
||||
}
|
||||
},
|
||||
]
|
||||
|
||||
result = calculate_vertex_ai_batch_cost_and_usage(
|
||||
responses, model_name="gemini-2.0-flash-001"
|
||||
)
|
||||
|
||||
assert result.usage.prompt_tokens == 18
|
||||
assert result.usage.completion_tokens == 8
|
||||
assert result.usage.total_tokens == 26
|
||||
assert result.cost > 0
|
||||
assert result.successful_requests == 2
|
||||
assert result.failed_requests == 1
|
||||
|
||||
def test_should_return_zeros_for_empty_response_list(self):
|
||||
"""Empty input → zero cost and zero usage."""
|
||||
|
|
@ -739,143 +646,4 @@ class TestVertexAIBatchCostCalculation:
|
|||
assert result.usage.completion_tokens == 0
|
||||
assert result.usage.total_tokens == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_shaped_output_records_nonzero_cost_and_usage(self):
|
||||
"""
|
||||
Regression test for the bug where Vertex batch cost/usage was always 0.
|
||||
|
||||
After PR #25627 (transform_file_content_response), the GCS predictions.jsonl
|
||||
is rewritten into OpenAI batch shape before the cost-tracking path sees it.
|
||||
With disable_vertex_batch_output_transformation=False (default), the cost
|
||||
dispatch must fall through to the generic aggregation path rather than
|
||||
calling calculate_vertex_ai_batch_cost_and_usage (which only reads raw
|
||||
usageMetadata fields).
|
||||
"""
|
||||
import litellm
|
||||
from litellm.batches.batch_utils import calculate_batch_cost_and_usage
|
||||
|
||||
openai_shaped_responses = [
|
||||
{
|
||||
"id": "batch_req_abc123",
|
||||
"custom_id": "request-1",
|
||||
"response": {
|
||||
"status_code": 200,
|
||||
"request_id": "chatcmpl-xyz",
|
||||
"body": {
|
||||
"id": "chatcmpl-xyz",
|
||||
"object": "chat.completion",
|
||||
"model": "gemini-2.0-flash-001",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Hello!"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
},
|
||||
},
|
||||
},
|
||||
"error": None,
|
||||
},
|
||||
{
|
||||
"id": "batch_req_def456",
|
||||
"custom_id": "request-2",
|
||||
"response": {
|
||||
"status_code": 200,
|
||||
"request_id": "chatcmpl-uvw",
|
||||
"body": {
|
||||
"id": "chatcmpl-uvw",
|
||||
"object": "chat.completion",
|
||||
"model": "gemini-2.0-flash-001",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "World!"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 8,
|
||||
"completion_tokens": 3,
|
||||
"total_tokens": 11,
|
||||
},
|
||||
},
|
||||
},
|
||||
"error": None,
|
||||
},
|
||||
]
|
||||
|
||||
original_flag = getattr(
|
||||
litellm, "disable_vertex_batch_output_transformation", False
|
||||
)
|
||||
try:
|
||||
litellm.disable_vertex_batch_output_transformation = False
|
||||
|
||||
result = await calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=openai_shaped_responses,
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_name="gemini-2.0-flash-001",
|
||||
)
|
||||
finally:
|
||||
litellm.disable_vertex_batch_output_transformation = original_flag
|
||||
|
||||
assert (
|
||||
result.usage.prompt_tokens == 18
|
||||
), f"expected 18 prompt tokens, got {result.usage.prompt_tokens}"
|
||||
assert (
|
||||
result.usage.completion_tokens == 8
|
||||
), f"expected 8 completion tokens, got {result.usage.completion_tokens}"
|
||||
assert (
|
||||
result.usage.total_tokens == 26
|
||||
), f"expected 26 total tokens, got {result.usage.total_tokens}"
|
||||
assert (
|
||||
result.cost > 0
|
||||
), f"expected non-zero cost for completed Vertex batch, got {result.cost}"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_vertex_output_still_works_when_transformation_disabled(self):
|
||||
"""
|
||||
When disable_vertex_batch_output_transformation=True the GCS file is returned
|
||||
as raw Vertex predictions.jsonl; the specialized reader must be used.
|
||||
"""
|
||||
import litellm
|
||||
from litellm.batches.batch_utils import calculate_batch_cost_and_usage
|
||||
|
||||
raw_vertex_responses = [
|
||||
{
|
||||
"request": {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]},
|
||||
"status": "",
|
||||
"response": {
|
||||
"candidates": [{"content": {"parts": [{"text": "Hello!"}]}}],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 10,
|
||||
"candidatesTokenCount": 5,
|
||||
"totalTokenCount": 15,
|
||||
},
|
||||
},
|
||||
"processed_time": "2026-01-01T00:00:00Z",
|
||||
},
|
||||
]
|
||||
|
||||
original_flag = getattr(
|
||||
litellm, "disable_vertex_batch_output_transformation", False
|
||||
)
|
||||
try:
|
||||
litellm.disable_vertex_batch_output_transformation = True
|
||||
|
||||
result = await calculate_batch_cost_and_usage(
|
||||
file_content_dictionary=raw_vertex_responses,
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_name="gemini-2.0-flash-001",
|
||||
)
|
||||
finally:
|
||||
litellm.disable_vertex_batch_output_transformation = original_flag
|
||||
|
||||
assert result.usage.prompt_tokens == 10
|
||||
assert result.usage.completion_tokens == 5
|
||||
assert result.usage.total_tokens == 15
|
||||
assert result.cost > 0, "raw Vertex shape should also produce non-zero cost"
|
||||
|
|
|
|||
|
|
@ -2198,6 +2198,34 @@ async def test_ProxyConfig_load_config_wires_general_settings_url_validation(tmp
|
|||
litellm.provider_url_destination_allowed_hosts = original_provider_hosts
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ssrf_block_message_names_a_config_section_load_config_honors(tmp_path, monkeypatch):
|
||||
"""Regression for LIT-8349: the remediation in the SSRF block message must point at a section that works."""
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
|
||||
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", [])
|
||||
monkeypatch.setattr(litellm, "user_url_validation", True)
|
||||
with pytest.raises(SSRFError) as blocked:
|
||||
validate_url("http://10.96.3.245:10002/agent.json")
|
||||
section_match = re.search(r"add the host to `user_url_allowed_hosts` in (\w+)\.", str(blocked.value))
|
||||
assert section_match is not None, str(blocked.value)
|
||||
section: Final = section_match.group(1)
|
||||
assert section == "litellm_settings", f"block message points admins at {section}, which the docs contradict"
|
||||
|
||||
f = tmp_path / "c.yaml"
|
||||
f.write_text(f"model_list: []\n{section}:\n user_url_allowed_hosts:\n - '10.96.3.245:10002'\n")
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
await ProxyConfig().load_config(router=None, config_file_path=str(f))
|
||||
|
||||
assert litellm.user_url_allowed_hosts == ["10.96.3.245:10002"], f"{section} did not apply the allowlist"
|
||||
assert validate_url("http://10.96.3.245:10002/agent.json") == (
|
||||
"http://10.96.3.245:10002/agent.json",
|
||||
"10.96.3.245:10002",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_wires_config_reload_interval(tmp_path, monkeypatch):
|
||||
"""general_settings.proxy_config_reload_interval_seconds must reach the proxy_server
|
||||
|
|
|
|||
|
|
@ -634,25 +634,6 @@ def test_equal_modeled_usage_is_zero_under_equivalent_model_names() -> None:
|
|||
assert _savings("claude-opus-5", "anthropic/claude-opus-5", usage, usage) == 0.0
|
||||
|
||||
|
||||
def test_baseline_is_priced_under_its_own_provider():
|
||||
"""Two providers can serve the same bare model name at different rates, so dropping
|
||||
the provider prices the baseline against a vendor the operator never named. Here it
|
||||
decides whether routing reads as a saving or a loss."""
|
||||
usage = Usage(prompt_tokens=100_000, completion_tokens=10_000, total_tokens=110_000)
|
||||
azure = compute_autorouter_savings(
|
||||
baseline_model="azure_ai/deepseek-r1",
|
||||
selected_model="claude-haiku-4-5",
|
||||
selected_provider="anthropic",
|
||||
usage=usage,
|
||||
)
|
||||
deepseek = compute_autorouter_savings(
|
||||
baseline_model="deepseek/deepseek-r1",
|
||||
selected_model="claude-haiku-4-5",
|
||||
selected_provider="anthropic",
|
||||
usage=usage,
|
||||
)
|
||||
assert azure != pytest.approx(deepseek)
|
||||
assert azure > 0 > deepseek
|
||||
|
||||
|
||||
def test_unresolvable_baseline_remains_unknown():
|
||||
|
|
|
|||
|
|
@ -3887,179 +3887,7 @@ class TestSpendLogsPayload:
|
|||
}
|
||||
return mock_response
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_logs_payload_success_log_with_api_base(self, monkeypatch):
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
# Clear any env overrides that would change the recorded api_base
|
||||
monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False)
|
||||
monkeypatch.delenv("ANTHROPIC_API_BASE", raising=False)
|
||||
|
||||
litellm.callbacks = [_ProxyDBLogger(message_logging=False)]
|
||||
# litellm._turn_on_debug()
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter,
|
||||
"_insert_spend_log_to_db",
|
||||
) as mock_client,
|
||||
patch.object(litellm.proxy.proxy_server, "prisma_client"),
|
||||
patch.object(client, "post", side_effect=self.mock_anthropic_response),
|
||||
):
|
||||
response = await litellm.acompletion(
|
||||
model="claude-4-sonnet-20250514",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
metadata={"user_api_key_end_user_id": "test_user_1"},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "Hi! My name is Claude."
|
||||
|
||||
await _wait_for_mock_call(mock_client)
|
||||
|
||||
kwargs = mock_client.call_args.kwargs
|
||||
payload: SpendLogsPayload = kwargs["payload"]
|
||||
expected_payload = SpendLogsPayload(
|
||||
**{
|
||||
"request_id": "chatcmpl-34df56d5-4807-45c1-bb99-61e52586b802",
|
||||
"call_type": "acompletion",
|
||||
"api_key": "",
|
||||
"cache_hit": "None",
|
||||
"startTime": datetime.datetime(
|
||||
2025, 3, 24, 22, 2, 42, 975883, tzinfo=datetime.timezone.utc
|
||||
),
|
||||
"endTime": datetime.datetime(
|
||||
2025, 3, 24, 22, 2, 42, 989132, tzinfo=datetime.timezone.utc
|
||||
),
|
||||
"completionStartTime": datetime.datetime(
|
||||
2025, 3, 24, 22, 2, 42, 989132, tzinfo=datetime.timezone.utc
|
||||
),
|
||||
"model": "claude-4-sonnet-20250514",
|
||||
"user": "",
|
||||
"team_id": "",
|
||||
"metadata": '{"applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-4-sonnet-20250514", "model_map_value": {"key": "claude-4-sonnet-20250514", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}',
|
||||
"cache_key": "Cache OFF",
|
||||
"spend": 0.01383,
|
||||
"total_tokens": 2598,
|
||||
"prompt_tokens": 2095,
|
||||
"completion_tokens": 503,
|
||||
"request_tags": "[]",
|
||||
"end_user": "test_user_1",
|
||||
"api_base": "https://api.anthropic.com/v1/messages",
|
||||
"model_group": "",
|
||||
"model_id": "",
|
||||
"requester_ip_address": None,
|
||||
"custom_llm_provider": "anthropic",
|
||||
"messages": "{}",
|
||||
"response": "{}",
|
||||
"proxy_server_request": "{}",
|
||||
"status": "success",
|
||||
"mcp_namespaced_tool_name": None,
|
||||
"agent_id": None,
|
||||
}
|
||||
)
|
||||
|
||||
differences = _compare_nested_dicts(
|
||||
payload, expected_payload, ignore_keys=ignored_keys
|
||||
)
|
||||
if differences:
|
||||
pytest.fail(f"Dictionary mismatch: {differences}")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_logs_payload_success_log_with_router(self, monkeypatch):
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
# Clear any env overrides that would change the recorded api_base
|
||||
monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False)
|
||||
monkeypatch.delenv("ANTHROPIC_API_BASE", raising=False)
|
||||
|
||||
litellm.callbacks = [_ProxyDBLogger(message_logging=False)]
|
||||
# litellm._turn_on_debug()
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "my-anthropic-model-group",
|
||||
"litellm_params": {
|
||||
"model": "claude-4-sonnet-20250514",
|
||||
},
|
||||
"model_info": {
|
||||
"id": "my-unique-model-id",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter,
|
||||
"_insert_spend_log_to_db",
|
||||
) as mock_client,
|
||||
patch.object(litellm.proxy.proxy_server, "prisma_client"),
|
||||
patch.object(client, "post", side_effect=self.mock_anthropic_response),
|
||||
):
|
||||
response = await router.acompletion(
|
||||
model="my-anthropic-model-group",
|
||||
messages=[{"role": "user", "content": "Hello, world!"}],
|
||||
metadata={"user_api_key_end_user_id": "test_user_1"},
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "Hi! My name is Claude."
|
||||
|
||||
await _wait_for_mock_call(mock_client)
|
||||
|
||||
kwargs = mock_client.call_args.kwargs
|
||||
payload: SpendLogsPayload = kwargs["payload"]
|
||||
expected_payload = SpendLogsPayload(
|
||||
**{
|
||||
"request_id": "chatcmpl-34df56d5-4807-45c1-bb99-61e52586b802",
|
||||
"call_type": "acompletion",
|
||||
"api_key": "",
|
||||
"cache_hit": "None",
|
||||
"startTime": datetime.datetime(
|
||||
2025, 3, 24, 22, 2, 42, 975883, tzinfo=datetime.timezone.utc
|
||||
),
|
||||
"endTime": datetime.datetime(
|
||||
2025, 3, 24, 22, 2, 42, 989132, tzinfo=datetime.timezone.utc
|
||||
),
|
||||
"completionStartTime": datetime.datetime(
|
||||
2025, 3, 24, 22, 2, 42, 989132, tzinfo=datetime.timezone.utc
|
||||
),
|
||||
"model": "claude-4-sonnet-20250514",
|
||||
"user": "",
|
||||
"team_id": "",
|
||||
"metadata": '{"applied_guardrails": [], "attempted_fallbacks": 0, "original_model_group": "my-anthropic-model-group", "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 503, "prompt_tokens": 2095, "total_tokens": 2598, "completion_tokens_details": null, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}, "model_map_information": {"model_map_key": "claude-4-sonnet-20250514", "model_map_value": {"key": "claude-4-sonnet-20250514", "max_tokens": 128000, "max_input_tokens": 200000, "max_output_tokens": 128000, "input_cost_per_token": 3e-06, "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": null, "output_cost_per_token_batches": null, "output_cost_per_token": 1.5e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "anthropic", "mode": "chat", "supports_system_messages": null, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": true, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": true, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": false, "supports_reasoning": true, "search_context_cost_per_query": null, "tpm": null, "rpm": null, "supported_openai_params": ["stream", "stop", "temperature", "top_p", "max_tokens", "max_completion_tokens", "tools", "tool_choice", "extra_headers", "parallel_tool_calls", "response_format", "user", "reasoning_effort", "thinking"]}}, "additional_usage_values": {"completion_tokens_details": {"accepted_prediction_tokens": null, "audio_tokens": null, "reasoning_tokens": null, "rejected_prediction_tokens": null, "text_tokens": 503, "image_tokens": null}, "prompt_tokens_details": {"audio_tokens": null, "cached_tokens": 0, "text_tokens": null, "image_tokens": null}, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0}}',
|
||||
"cache_key": "Cache OFF",
|
||||
"spend": 0.01383,
|
||||
"total_tokens": 2598,
|
||||
"prompt_tokens": 2095,
|
||||
"completion_tokens": 503,
|
||||
"request_tags": "[]",
|
||||
"end_user": "test_user_1",
|
||||
"api_base": "https://api.anthropic.com/v1/messages",
|
||||
"model_group": "my-anthropic-model-group",
|
||||
"model_id": "my-unique-model-id",
|
||||
"requester_ip_address": None,
|
||||
"custom_llm_provider": "anthropic",
|
||||
"messages": "{}",
|
||||
"response": "{}",
|
||||
"proxy_server_request": "{}",
|
||||
"status": "success",
|
||||
"mcp_namespaced_tool_name": None,
|
||||
"agent_id": None,
|
||||
}
|
||||
)
|
||||
|
||||
differences = _compare_nested_dicts(
|
||||
payload, expected_payload, ignore_keys=ignored_keys
|
||||
)
|
||||
if differences:
|
||||
pytest.fail(f"Dictionary mismatch: {differences}")
|
||||
|
||||
|
||||
def _compare_nested_dicts(
|
||||
|
|
|
|||
|
|
@ -8238,3 +8238,45 @@ def test_default_team_settings_bool_turn_off_message_logging_redacts():
|
|||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ["/mcp-rest/tools/call", "/v1/responses", "/v1/chat/completions"])
|
||||
@pytest.mark.parametrize("custom_auth", ["x-mcp-auth", "x-private-mcp-token"])
|
||||
async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custom_auth: str):
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
metadata_name: Final = "litellm_metadata" if path == "/v1/responses" else "metadata"
|
||||
secrets: Final = {
|
||||
"X-MCP-Deepwiki-Authorization": "upstream-sentinel",
|
||||
custom_auth: "client-auth-sentinel",
|
||||
"x-service-token": "configured-secret-sentinel",
|
||||
}
|
||||
attribution: Final = {"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a"}
|
||||
request: Final = _make_request_mock(path, {"Content-Type": "application/json", **secrets, **attribution})
|
||||
request.headers = Headers(request.headers)
|
||||
settings: Final = {"mcp_client_side_auth_header_name": custom_auth, "user_header_name": "x-user-id"}
|
||||
server: Final = MCPServer(
|
||||
server_id="header-test", name="header-test", transport="http", url="https://example.com/mcp",
|
||||
extra_headers=["x-service-token", "x-user-id"],
|
||||
)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.general_settings", settings),
|
||||
patch.dict(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.config_mcp_servers",
|
||||
{"header-test": server}, clear=True,
|
||||
),
|
||||
):
|
||||
updated: Final = await add_litellm_data_to_request(
|
||||
data={"model": "test-model", "messages": [{"role": "user", "content": "hello"}]},
|
||||
request=request, user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
|
||||
proxy_config=MagicMock(), general_settings=settings, version="test",
|
||||
)
|
||||
for header_dict in _all_header_dicts(updated, metadata_name):
|
||||
assert not any(value in json.dumps(header_dict) for value in secrets.values())
|
||||
assert updated[metadata_name]["headers"] == updated["proxy_server_request"]["headers"]
|
||||
for name, value in attribution.items():
|
||||
assert updated[metadata_name]["headers"][name] == value
|
||||
for name, value in secrets.items():
|
||||
assert updated["secret_fields"]["raw_headers"][name.lower()] == value
|
||||
assert request.headers[name] == value
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import pytest
|
|||
from fastapi import FastAPI, Request
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy
|
||||
|
|
@ -112,8 +113,10 @@ async def test_real_proxy_child_auth_privacy_and_body_policy(
|
|||
})))
|
||||
return asyncio.sleep(0, result=ModelResponse(id="private-summary", model="compactor"))
|
||||
|
||||
monkeypatch.setattr(litellm, "max_budget", 0)
|
||||
monkeypatch.setattr(proxy_server.app, "dependency_overrides", {})
|
||||
monkeypatch.setattr(proxy_server, "master_key", "sk-master-fixture")
|
||||
monkeypatch.setattr(litellm, "max_budget", 0.0)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", object())
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", None)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,8 @@ import pytest
|
|||
|
||||
|
||||
import builtins
|
||||
import runpy
|
||||
import sys
|
||||
import types
|
||||
import urllib.parse as urlparse
|
||||
|
||||
|
|
@ -18,6 +20,7 @@ import yaml
|
|||
from uvicorn.config import LOOP_FACTORIES
|
||||
from uvicorn.importer import import_from_string
|
||||
|
||||
from litellm.proxy import proxy_cli
|
||||
from litellm.proxy.proxy_cli import ProxyInitializationHelpers, run_server
|
||||
|
||||
|
||||
|
|
@ -636,6 +639,36 @@ class TestProxyInitializationHelpers:
|
|||
), f"exit_code={result.exit_code}, output={result.output}"
|
||||
mock_uvicorn_run.assert_called_once()
|
||||
|
||||
@patch("uvicorn.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
@patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False)
|
||||
def test_script_boot_imports_the_package_proxy_server(
|
||||
self, mock_should_update, mock_setup_db, mock_atexit_register, mock_uvicorn_run
|
||||
):
|
||||
package_proxy_server = MagicMock(
|
||||
app=MagicMock(),
|
||||
ProxyConfig=MagicMock(),
|
||||
KeyManagementSettings=MagicMock(),
|
||||
save_worker_config=MagicMock(),
|
||||
)
|
||||
sibling_proxy_server = types.ModuleType("proxy_server")
|
||||
clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")}
|
||||
with (
|
||||
patch.dict(os.environ, clean_env, clear=True),
|
||||
patch.dict(
|
||||
"sys.modules",
|
||||
{"proxy_server": sibling_proxy_server, "litellm.proxy.proxy_server": package_proxy_server},
|
||||
),
|
||||
patch.object(sys, "argv", ["proxy_cli.py", "--skip_server_startup"]),
|
||||
patch.object(sys, "path", list(sys.path)),
|
||||
pytest.raises(SystemExit) as exit_info,
|
||||
):
|
||||
runpy.run_path(proxy_cli.__file__, run_name="__main__")
|
||||
|
||||
assert exit_info.value.code == 0
|
||||
package_proxy_server.save_worker_config.assert_called_once()
|
||||
|
||||
@patch("uvicorn.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
|
|
@ -1174,6 +1207,172 @@ class TestProxyInitializationHelpers:
|
|||
else:
|
||||
assert "pgbouncer" not in appended_params
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"env_value, config_value, expect_pgbouncer",
|
||||
[
|
||||
("true", None, True),
|
||||
("1", None, True),
|
||||
("false", None, False),
|
||||
(None, None, False),
|
||||
("true", False, True),
|
||||
("false", True, True),
|
||||
],
|
||||
)
|
||||
@patch("subprocess.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
@patch(
|
||||
"litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False
|
||||
)
|
||||
def test_disable_prepared_statements_env_var_forwarded_to_url(
|
||||
self,
|
||||
mock_should_update,
|
||||
mock_setup_db,
|
||||
mock_atexit_register,
|
||||
mock_subprocess_run,
|
||||
env_value,
|
||||
config_value,
|
||||
expect_pgbouncer,
|
||||
):
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
runner = CliRunner()
|
||||
mock_subprocess_run.return_value = MagicMock(returncode=0)
|
||||
|
||||
general_settings = {"database_url": "postgresql://test:test@localhost:5432/test"}
|
||||
if config_value is not None:
|
||||
general_settings["database_disable_prepared_statements"] = config_value
|
||||
mock_proxy_module = MagicMock(
|
||||
app=MagicMock(),
|
||||
ProxyConfig=MagicMock(),
|
||||
KeyManagementSettings=MagicMock(),
|
||||
save_worker_config=MagicMock(),
|
||||
)
|
||||
mock_proxy_module.ProxyConfig.return_value.get_config = AsyncMock(
|
||||
return_value={"general_settings": general_settings}
|
||||
)
|
||||
|
||||
clean_env = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in ("DATABASE_URL", "DIRECT_URL", "DATABASE_DISABLE_PREPARED_STATEMENTS")
|
||||
}
|
||||
if env_value is not None:
|
||||
clean_env["DATABASE_DISABLE_PREPARED_STATEMENTS"] = env_value
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, clean_env, clear=True),
|
||||
patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"proxy_server": mock_proxy_module,
|
||||
"litellm.proxy.proxy_server": mock_proxy_module,
|
||||
},
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
|
||||
) as mock_get_args,
|
||||
patch(
|
||||
"litellm.proxy.proxy_cli.append_query_params",
|
||||
side_effect=lambda url, params: str(url),
|
||||
) as mock_append_query_params,
|
||||
):
|
||||
mock_get_args.return_value = {
|
||||
"app": "litellm.proxy.proxy_server:app",
|
||||
"host": "localhost",
|
||||
"port": 8000,
|
||||
}
|
||||
|
||||
result = runner.invoke(
|
||||
run_server,
|
||||
["--local", "--config", "test-config.yaml", "--skip_server_startup"],
|
||||
)
|
||||
|
||||
assert (
|
||||
result.exit_code == 0
|
||||
), f"exit_code={result.exit_code}, output={result.output}"
|
||||
appended_params = mock_append_query_params.call_args.args[1]
|
||||
if expect_pgbouncer:
|
||||
assert appended_params["pgbouncer"] == "true", appended_params
|
||||
else:
|
||||
assert "pgbouncer" not in appended_params, appended_params
|
||||
|
||||
@patch("subprocess.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
@patch(
|
||||
"litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False
|
||||
)
|
||||
def test_malformed_disable_prepared_statements_env_var_is_rejected_even_when_config_enables_it(
|
||||
self,
|
||||
mock_should_update,
|
||||
mock_setup_db,
|
||||
mock_atexit_register,
|
||||
mock_subprocess_run,
|
||||
):
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
runner = CliRunner()
|
||||
mock_subprocess_run.return_value = MagicMock(returncode=0)
|
||||
|
||||
mock_proxy_module = MagicMock(
|
||||
app=MagicMock(),
|
||||
ProxyConfig=MagicMock(),
|
||||
KeyManagementSettings=MagicMock(),
|
||||
save_worker_config=MagicMock(),
|
||||
)
|
||||
mock_proxy_module.ProxyConfig.return_value.get_config = AsyncMock(
|
||||
return_value={
|
||||
"general_settings": {
|
||||
"database_url": "postgresql://test:test@localhost:5432/test",
|
||||
"database_disable_prepared_statements": True,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
clean_env = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in ("DATABASE_URL", "DIRECT_URL", "DATABASE_DISABLE_PREPARED_STATEMENTS")
|
||||
}
|
||||
clean_env["DATABASE_DISABLE_PREPARED_STATEMENTS"] = "enabled"
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, clean_env, clear=True),
|
||||
patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"proxy_server": mock_proxy_module,
|
||||
"litellm.proxy.proxy_server": mock_proxy_module,
|
||||
},
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
|
||||
) as mock_get_args,
|
||||
patch(
|
||||
"litellm.proxy.proxy_cli.append_query_params",
|
||||
side_effect=lambda url, params: str(url),
|
||||
) as mock_append_query_params,
|
||||
):
|
||||
mock_get_args.return_value = {
|
||||
"app": "litellm.proxy.proxy_server:app",
|
||||
"host": "localhost",
|
||||
"port": 8000,
|
||||
}
|
||||
|
||||
result = runner.invoke(
|
||||
run_server,
|
||||
["--local", "--config", "test-config.yaml", "--skip_server_startup"],
|
||||
)
|
||||
|
||||
assert isinstance(result.exception, ValueError), f"exit_code={result.exit_code}, output={result.output}"
|
||||
assert "DATABASE_DISABLE_PREPARED_STATEMENTS" in str(result.exception), result.exception
|
||||
mock_append_query_params.assert_not_called()
|
||||
|
||||
@patch("uvicorn.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
|
|
@ -1779,13 +1978,18 @@ class TestProxyInitializationHelpers:
|
|||
mock_proxy_config_instance.get_config = mock_get_config
|
||||
mock_proxy_config.return_value = mock_proxy_config_instance
|
||||
|
||||
mock_proxy_server_module = MagicMock(app=mock_app)
|
||||
mock_proxy_server_module = MagicMock(
|
||||
app=mock_app,
|
||||
ProxyConfig=mock_proxy_config,
|
||||
KeyManagementSettings=mock_key_mgmt,
|
||||
save_worker_config=mock_save_worker_config,
|
||||
)
|
||||
|
||||
# Only remove DATABASE_URL and DIRECT_URL to prevent the database setup
|
||||
# code path from running. Do NOT use clear=True as it removes PATH, HOME,
|
||||
# etc., which causes imports inside run_server to break in CI (the real
|
||||
# litellm.proxy.proxy_server import at line 820 of proxy_cli.py has heavy
|
||||
# side effects that fail without a proper environment).
|
||||
# litellm.proxy.proxy_server import has heavy side effects that fail
|
||||
# without a proper environment).
|
||||
env_overrides = {
|
||||
"DATABASE_URL": "",
|
||||
"DIRECT_URL": "",
|
||||
|
|
@ -1801,18 +2005,7 @@ class TestProxyInitializationHelpers:
|
|||
with (
|
||||
patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"proxy_server": MagicMock(
|
||||
app=mock_app,
|
||||
ProxyConfig=mock_proxy_config,
|
||||
KeyManagementSettings=mock_key_mgmt,
|
||||
save_worker_config=mock_save_worker_config,
|
||||
),
|
||||
# Also mock litellm.proxy.proxy_server to prevent the real
|
||||
# import at line 820 of proxy_cli.py which has heavy side
|
||||
# effects (FastAPI app init, logging setup, etc.)
|
||||
"litellm.proxy.proxy_server": mock_proxy_server_module,
|
||||
},
|
||||
{"litellm.proxy.proxy_server": mock_proxy_server_module},
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
|
||||
|
|
|
|||
|
|
@ -3200,7 +3200,7 @@ def test_normalize_datetime_for_sorting():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_proxy_budget_to_db_only_creates_user_no_keys():
|
||||
async def test_add_proxy_budget_to_db_only_creates_user_no_keys(monkeypatch: pytest.MonkeyPatch):
|
||||
"""
|
||||
Test that _add_proxy_budget_to_db only creates a user and no keys are added.
|
||||
|
||||
|
|
@ -3218,8 +3218,8 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys():
|
|||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
|
||||
# Set up required litellm settings
|
||||
litellm.budget_duration = "30d"
|
||||
litellm.max_budget = 100.0
|
||||
monkeypatch.setattr(litellm, "budget_duration", "30d")
|
||||
monkeypatch.setattr(litellm, "max_budget", 100.0)
|
||||
|
||||
litellm_proxy_budget_name = "litellm-proxy-budget"
|
||||
|
||||
|
|
@ -3258,7 +3258,7 @@ async def test_add_proxy_budget_to_db_only_creates_user_no_keys():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_proxy_budget_to_db_backfills_budget_reset_at():
|
||||
async def test_add_proxy_budget_to_db_backfills_budget_reset_at(monkeypatch: pytest.MonkeyPatch):
|
||||
"""
|
||||
Test that _upsert_proxy_budget_with_reset_at_backfill issues a conditional
|
||||
update_many with `WHERE budget_reset_at IS NULL` to backfill the column on
|
||||
|
|
@ -3276,8 +3276,8 @@ async def test_add_proxy_budget_to_db_backfills_budget_reset_at():
|
|||
import litellm
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
|
||||
litellm.budget_duration = "30d"
|
||||
litellm.max_budget = 100.0
|
||||
monkeypatch.setattr(litellm, "budget_duration", "30d")
|
||||
monkeypatch.setattr(litellm, "max_budget", 100.0)
|
||||
litellm_proxy_budget_name = "litellm-proxy-budget"
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
|
|
|
|||
|
|
@ -5,16 +5,18 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import Any, Final
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm.constants import PROXY_REJECTED_BEFORE_ROUTING_KEY
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import AlertType, ProxyErrorTypes, UserAPIKeyAuth
|
||||
from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
|
||||
|
|
@ -99,6 +101,456 @@ async def test_post_call_failure_hook_no_callbacks_returns_none(
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_attributes_single_router_deployment(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
"model_info": {"provider": "acme"},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={"model": "internal-model", "messages": [{"role": "user", "content": "hi"}]},
|
||||
original_exception=HTTPException(status_code=403, detail="blocked"),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs["custom_llm_provider"] == "openai"
|
||||
assert kwargs["litellm_params"]["custom_llm_provider"] == "openai"
|
||||
assert kwargs["litellm_params"]["metadata"]["model_info"]["provider"] == "acme"
|
||||
assert kwargs["litellm_params"]["metadata"]["deployment"] == "openai/gpt-4.1"
|
||||
assert kwargs["litellm_params"][PROXY_REJECTED_BEFORE_ROUTING_KEY] is True
|
||||
assert kwargs["standard_logging_object"]["custom_llm_provider"] == "openai"
|
||||
assert (
|
||||
kwargs["standard_logging_object"]["model_id"] == proxy_server.llm_router.get_model_list()[0]["model_info"]["id"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_keeps_router_stamped_metadata_for_post_call_failures(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""A post-call guardrail block arrives after the provider handoff with the router's own
|
||||
``model_info`` in the request metadata. The pre-routing flag must stay off so deployment
|
||||
metrics keep attributing the failure to the deployment that actually served the call."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
"model_info": {"id": "routed-deployment"},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
request_data = {
|
||||
"litellm_call_id": "post-call-guardrail",
|
||||
"model": "internal-model",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {"model_info": {"id": "routed-deployment", "served": True}},
|
||||
}
|
||||
logging_obj, request_data = litellm.utils.function_setup(
|
||||
original_function="acompletion", rules_obj=litellm.utils.Rules(), start_time=datetime.now(), **request_data
|
||||
)
|
||||
logging_obj.model_call_details["first_api_call_start_time"] = datetime.now()
|
||||
request_data["litellm_logging_obj"] = logging_obj
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data=request_data,
|
||||
original_exception=GuardrailRaisedException(guardrail_name="g", message="response blocked"),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs["litellm_params"]["metadata"]["model_info"] == {"id": "routed-deployment", "served": True}
|
||||
assert PROXY_REJECTED_BEFORE_ROUTING_KEY not in kwargs["litellm_params"]
|
||||
assert kwargs["standard_logging_object"]["model_id"] == "routed-deployment"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_flags_pre_routing_reject_despite_caller_model_info(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""A key allowed to override pricing keeps caller-supplied ``metadata.model_info``. A reject
|
||||
before any provider handoff must still carry the pre-routing flag so deployment metrics do
|
||||
not record an outage for a deployment the request never reached."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
"model_info": {"id": "real-deployment"},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={
|
||||
"model": "internal-model",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {"model_info": {"id": "spoofed-deployment"}},
|
||||
},
|
||||
original_exception=HTTPException(status_code=429, detail="key over limit"),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs["litellm_params"][PROXY_REJECTED_BEFORE_ROUTING_KEY] is True
|
||||
assert kwargs["litellm_params"]["custom_llm_provider"] == "openai"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_attribution_does_not_count_against_the_deployment(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""The router's failure callbacks run on this path too. A proxy-side reject must not
|
||||
bump the deployment's failure or rpm counters, or a key hitting its own limit
|
||||
could cool down the only deployment for everyone."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test", "rpm": 100},
|
||||
}
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
proxy_logging.alert_types = []
|
||||
deployment_id = router.get_model_list()[0]["model_info"]["id"]
|
||||
|
||||
for status in (403, 429):
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={"model": "internal-model", "messages": [{"role": "user", "content": "hi"}]},
|
||||
original_exception=HTTPException(status_code=status, detail="blocked"),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
pending = asyncio.all_tasks() - {asyncio.current_task()}
|
||||
await asyncio.gather(*pending, return_exceptions=True)
|
||||
|
||||
deployment_keys = [key for key in router.cache.in_memory_cache.cache_dict if deployment_id in key]
|
||||
assert deployment_keys == [], f"proxy reject was counted against the deployment: {deployment_keys}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_attributes_the_keys_team_deployment_over_the_global_group(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""A team key requesting its team public model name must be attributed to the team's
|
||||
deployment, not to a global group that happens to share the public name."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "shared-name",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
"model_info": {"id": "global-deployment"},
|
||||
},
|
||||
{
|
||||
"model_name": "shared-name_test-team_deadbeef",
|
||||
"litellm_params": {"model": "anthropic/claude-sonnet-4-5", "api_key": "sk-test"},
|
||||
"model_info": {
|
||||
"id": "team-deployment",
|
||||
"team_id": "test-team",
|
||||
"team_public_model_name": "shared-name",
|
||||
},
|
||||
},
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={"model": "shared-name", "messages": [{"role": "user", "content": "hi"}]},
|
||||
original_exception=HTTPException(status_code=429, detail="rate limited"),
|
||||
user_api_key_dict=make_user_api_key_auth(team_id="test-team", request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs["custom_llm_provider"] == "anthropic"
|
||||
assert kwargs["litellm_params"]["metadata"]["deployment"] == "anthropic/claude-sonnet-4-5"
|
||||
assert kwargs["standard_logging_object"]["model_id"] == "team-deployment"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_omits_provider_for_mixed_router_deployments(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
},
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "anthropic/claude-sonnet-4-5", "api_key": "sk-test"},
|
||||
},
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={"model": "internal-model", "messages": [{"role": "user", "content": "hi"}]},
|
||||
original_exception=HTTPException(status_code=403, detail="blocked"),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs.get("custom_llm_provider") is None
|
||||
assert "model_info" not in (kwargs["litellm_params"].get("metadata") or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_omits_provider_when_a_deployment_does_not_resolve(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""One deployment resolves to openai and its sibling resolves to nothing: the group
|
||||
is not known to be single-provider, so no provider is stamped on the failure."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"model_name": "internal-model", "litellm_params": {"model": "openai/gpt-4.1"}},
|
||||
{"model_name": "internal-model", "litellm_params": {"model": "unmapped-model-with-no-provider"}},
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={"model": "internal-model", "messages": [{"role": "user", "content": "hi"}]},
|
||||
original_exception=HTTPException(status_code=403, detail="blocked"),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs.get("custom_llm_provider") is None
|
||||
assert kwargs["litellm_params"].get("custom_llm_provider") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_logging_proxy_only_path_attributes_with_read_only_metadata(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
"""With a logging object already on the request, its metadata is taken as given;
|
||||
a read-only mapping there must not crash the stamp, and the failure handler
|
||||
still receives the provider attribution."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
"model_info": {"provider": "acme"},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.call_type = "acompletion"
|
||||
logging_obj.model_call_details = {}
|
||||
logging_obj.async_failure_handler = AsyncMock()
|
||||
|
||||
await proxy_logging._handle_logging_proxy_only_error(
|
||||
request_data={
|
||||
"litellm_logging_obj": logging_obj,
|
||||
"model": "internal-model",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": MappingProxyType({"user_api_key_alias": "frozen"}),
|
||||
},
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
original_exception=HTTPException(status_code=403, detail="blocked"),
|
||||
)
|
||||
|
||||
assert logging_obj.async_failure_handler.called
|
||||
update_kwargs = logging_obj.update_environment_variables.call_args.kwargs
|
||||
assert update_kwargs["custom_llm_provider"] == "openai"
|
||||
assert update_kwargs["litellm_params"]["custom_llm_provider"] == "openai"
|
||||
assert update_kwargs["litellm_params"]["metadata"] == {"user_api_key_alias": "frozen"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_fires_without_router_attribution(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
):
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "different-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={"model": "internal-model", "messages": [{"role": "user", "content": "hi"}]},
|
||||
original_exception=HTTPException(status_code=403, detail="blocked"),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs.get("custom_llm_provider") is None
|
||||
assert "model_info" not in (kwargs["litellm_params"].get("metadata") or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", [123, ["internal-model"], {"name": "internal-model"}, None])
|
||||
async def test_post_call_failure_hook_fires_for_non_string_model(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch, model: object
|
||||
):
|
||||
"""A body whose ``model`` is not a string is rejected by the proxy before routing; its
|
||||
failure callback must still fire, unattributed, instead of a TypeError escaping the hook."""
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
recorded: list[dict] = []
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
recorded.append(kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
proxy_server,
|
||||
"llm_router",
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "internal-model",
|
||||
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
|
||||
proxy_logging.alert_types = []
|
||||
|
||||
await proxy_logging.post_call_failure_hook(
|
||||
request_data={"model": model, "messages": [{"role": "user", "content": "hi"}]},
|
||||
original_exception=HTTPException(status_code=400, detail="'model' must be a string."),
|
||||
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert len(recorded) == 1
|
||||
kwargs = recorded[0]
|
||||
assert kwargs.get("custom_llm_provider") is None
|
||||
assert "model_info" not in (kwargs["litellm_params"].get("metadata") or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_failure_hook_callback_returns_http_exception(
|
||||
proxy_logging, make_user_api_key_auth, monkeypatch
|
||||
|
|
|
|||
|
|
@ -159,6 +159,7 @@ async def test_acompletion_with_mcp_passes_mcp_server_auth_headers_to_process_to
|
|||
secret_fields=secret_fields,
|
||||
)
|
||||
|
||||
assert captured_process_kwargs["raw_headers"] == secret_fields["raw_headers"]
|
||||
assert "mcp_server_auth_headers" in captured_process_kwargs
|
||||
mcp_server_auth_headers = captured_process_kwargs["mcp_server_auth_headers"]
|
||||
assert mcp_server_auth_headers is not None
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import sys
|
|||
import textwrap
|
||||
import types
|
||||
from contextlib import nullcontext
|
||||
from typing import Any, cast
|
||||
from typing import Any, Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -1235,6 +1235,7 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false(
|
|||
)
|
||||
|
||||
async def fake_process(**kwargs: Any) -> tuple[list[Any], dict[str, str]]:
|
||||
assert kwargs["raw_headers"] == {"x-app-id": "follow-up-caller"}
|
||||
return ([], {"foo": "litellm_proxy"})
|
||||
|
||||
async def fake_execute(**kwargs: Any) -> list[dict[str, Any]]:
|
||||
|
|
@ -1254,6 +1255,7 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false(
|
|||
model="gpt-5",
|
||||
tools=[{"type": "mcp", "server_url": "litellm_proxy", "require_approval": "never"}],
|
||||
litellm_metadata={"guardrails": ["block-all"]},
|
||||
secret_fields={"raw_headers": {"x-app-id": "follow-up-caller"}},
|
||||
store=store,
|
||||
previous_response_id=caller_previous_response_id,
|
||||
)
|
||||
|
|
@ -1266,3 +1268,40 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false(
|
|||
item for item in follow_up_call["input"] if isinstance(item, dict) and item.get("type") == "reasoning"
|
||||
]
|
||||
assert bool(reasoning_items) is (store is False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.proxy._experimental.mcp_server import operations
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
|
||||
headers: Final = {
|
||||
"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a",
|
||||
"x-mcp-deepwiki-authorization": "upstream-sentinel", "authorization": "proxy-sentinel",
|
||||
}
|
||||
manager: Final = types.SimpleNamespace(
|
||||
catalog=types.SimpleNamespace(operation=nullcontext),
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
)
|
||||
logger: Final = MagicMock(model_call_details={})
|
||||
logger.async_success_handler = AsyncMock()
|
||||
setup: Final = MagicMock(return_value=(logger, None))
|
||||
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[]))
|
||||
monkeypatch.setattr(operations, "function_setup", setup)
|
||||
response: Final = ResponsesAPIResponse(
|
||||
id="resp_test", created_at=1234567891, model="test-model", object="response",
|
||||
status="completed", output=[], parallel_tool_calls=False, tool_choice="auto", tools=[],
|
||||
)
|
||||
monkeypatch.setattr(responses_main, "aresponses", AsyncMock(return_value=response))
|
||||
result: Final = await responses_main.aresponses_api_with_mcp(
|
||||
input="hi", model="test-model", tools=[{"type": "mcp", "server_url": "litellm_proxy"}],
|
||||
secret_fields={"raw_headers": headers},
|
||||
)
|
||||
assert result is response
|
||||
logger.async_success_handler.assert_awaited_once()
|
||||
logged: Final = setup.call_args.kwargs["metadata"]["headers"]
|
||||
assert logged == {"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a"}
|
||||
assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel"
|
||||
|
|
|
|||
|
|
@ -48,72 +48,6 @@ class MetadataCaptureCallback(CustomLogger):
|
|||
self.event.set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_metadata_passed_to_custom_callback_codex_models():
|
||||
"""
|
||||
Test that metadata passed to completion() is available in custom callback
|
||||
when using codex models (responses API bridge path).
|
||||
|
||||
Codex models have mode=responses and route through responses_api_bridge,
|
||||
which passes litellm_metadata. The fix ensures this is preserved as
|
||||
litellm_params.metadata for callback compatibility.
|
||||
"""
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
mock_response = ResponsesAPIResponse.model_construct(
|
||||
id="resp-test",
|
||||
created_at=0,
|
||||
output=[
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg-1",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Hello!"}],
|
||||
}
|
||||
],
|
||||
object="response",
|
||||
model="gpt-5.1-codex",
|
||||
status="completed",
|
||||
usage={
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 10,
|
||||
"total_tokens": 15,
|
||||
},
|
||||
)
|
||||
|
||||
test_metadata = {"foo": "bar", "trace_id": "test-123"}
|
||||
callback = MetadataCaptureCallback()
|
||||
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
|
||||
litellm.callbacks = [callback]
|
||||
|
||||
try:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_post:
|
||||
mock_post.return_value = _make_mock_http_response(
|
||||
mock_response.model_dump()
|
||||
)
|
||||
# gpt-5.1-codex has mode=responses - routes through responses bridge
|
||||
await litellm.acompletion(
|
||||
model="gpt-5.1-codex",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
metadata=test_metadata,
|
||||
)
|
||||
|
||||
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
|
||||
|
||||
assert callback.captured_kwargs is not None, "Callback should have been invoked"
|
||||
|
||||
litellm_params = callback.captured_kwargs.get("litellm_params", {})
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
|
||||
assert "foo" in metadata, "metadata['foo'] should be accessible in callback"
|
||||
assert metadata["foo"] == "bar"
|
||||
assert metadata.get("trace_id") == "test-123"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -148,7 +148,6 @@ async def test_each_tokenizer_gets_its_own_cached_counter(fake_tokenizers: None)
|
|||
("gpt-4o", "o200k_base"),
|
||||
("gpt-4o-mini", "o200k_base"),
|
||||
("gpt-4o-2024-08-06", "o200k_base"),
|
||||
("chatgpt-4o-latest", "o200k_base"),
|
||||
("gpt-4.1", "o200k_base"),
|
||||
("gpt-5", "o200k_base"),
|
||||
("gpt-5-mini", "o200k_base"),
|
||||
|
|
|
|||
|
|
@ -159,131 +159,8 @@ def test_cost_calculator_with_response_cost_in_additional_headers():
|
|||
assert result == 1000
|
||||
|
||||
|
||||
def test_cost_calculator_with_usage(_local_model_cost_map, monkeypatch):
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=120,
|
||||
completion_tokens=100,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=10,
|
||||
audio_tokens=90,
|
||||
image_tokens=20,
|
||||
),
|
||||
)
|
||||
mr = ModelResponse(usage=usage, model="gemini-2.0-flash-001")
|
||||
|
||||
result = response_cost_calculator(
|
||||
response_object=mr,
|
||||
model="",
|
||||
custom_llm_provider="vertex_ai",
|
||||
call_type="acompletion",
|
||||
optional_params={},
|
||||
cache_hit=None,
|
||||
base_model=None,
|
||||
)
|
||||
|
||||
model_info = litellm.model_cost["gemini-2.0-flash-001"]
|
||||
|
||||
# Step 1: Test a model where input_cost_per_image_token is not set.
|
||||
# In this case the calculation should use input_cost_per_token as fallback.
|
||||
assert model_info.get("input_cost_per_image_token") is None, (
|
||||
"Test case expects that input_cost_per_image_token is not set"
|
||||
)
|
||||
|
||||
expected_cost = (
|
||||
usage.prompt_tokens_details.audio_tokens * model_info["input_cost_per_audio_token"]
|
||||
+ usage.prompt_tokens_details.text_tokens * model_info["input_cost_per_token"]
|
||||
+ usage.prompt_tokens_details.image_tokens * model_info["input_cost_per_token"]
|
||||
+ usage.completion_tokens * model_info["output_cost_per_token"]
|
||||
)
|
||||
|
||||
assert result == expected_cost, f"Got {result}, Expected {expected_cost}"
|
||||
|
||||
# Step 2: Set input_cost_per_image_token.
|
||||
# In this case the explicit cost information should be used.
|
||||
temp_model_info_object = dict(model_info)
|
||||
temp_model_info_object["input_cost_per_image_token"] = 0.5
|
||||
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"model_cost",
|
||||
{"gemini-2.0-flash-001": temp_model_info_object},
|
||||
)
|
||||
|
||||
# Invalidate caches after modifying litellm.model_cost
|
||||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||||
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
|
||||
result = response_cost_calculator(
|
||||
response_object=mr,
|
||||
model="",
|
||||
custom_llm_provider="vertex_ai",
|
||||
call_type="acompletion",
|
||||
optional_params={},
|
||||
cache_hit=None,
|
||||
base_model=None,
|
||||
)
|
||||
|
||||
expected_cost = (
|
||||
usage.prompt_tokens_details.audio_tokens * temp_model_info_object["input_cost_per_audio_token"]
|
||||
+ usage.prompt_tokens_details.text_tokens * temp_model_info_object["input_cost_per_token"]
|
||||
+ usage.prompt_tokens_details.image_tokens * temp_model_info_object["input_cost_per_image_token"]
|
||||
+ usage.completion_tokens * temp_model_info_object["output_cost_per_token"]
|
||||
)
|
||||
|
||||
assert result == expected_cost, f"Got {result}, Expected {expected_cost}"
|
||||
|
||||
|
||||
def test_handle_realtime_stream_cost_calculation_stores_cost_breakdown():
|
||||
"""Regression: realtime cost must populate logging_obj.cost_breakdown so the
|
||||
spend logs / UI show input vs output cost (issue: cost_breakdown was None for
|
||||
/v1/realtime even though a total spend was computed)."""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
results: OpenAIRealtimeStreamList = [
|
||||
{"type": "session.created", "session": {"model": "gpt-4o-realtime-preview"}},
|
||||
{
|
||||
"type": "response.done",
|
||||
"response": {
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"total_tokens": 150,
|
||||
}
|
||||
},
|
||||
},
|
||||
]
|
||||
combined_usage_object = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
|
||||
results=results,
|
||||
)
|
||||
|
||||
logging_obj = Logging(
|
||||
model="gpt-4o-realtime-preview",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="_arealtime",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="realtime-cost-breakdown-test",
|
||||
function_id="realtime-cost-breakdown-test",
|
||||
)
|
||||
|
||||
total_cost = handle_realtime_stream_cost_calculation(
|
||||
results=results,
|
||||
combined_usage_object=combined_usage_object,
|
||||
custom_llm_provider="openai",
|
||||
litellm_model_name="gpt-4o-realtime-preview",
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert total_cost > 0
|
||||
assert logging_obj.cost_breakdown is not None
|
||||
assert logging_obj.cost_breakdown["input_cost"] > 0
|
||||
assert logging_obj.cost_breakdown["output_cost"] > 0
|
||||
assert abs(logging_obj.cost_breakdown["input_cost"] + logging_obj.cost_breakdown["output_cost"] - total_cost) < 1e-9
|
||||
assert abs(logging_obj.cost_breakdown["total_cost"] - total_cost) < 1e-9
|
||||
|
||||
|
||||
def test_realtime_stream_combines_text_and_audio_token_details():
|
||||
|
|
@ -1124,126 +1001,6 @@ def test_bedrock_cost_calculator_comparison_with_without_cache():
|
|||
print(f"Cost with cache: {cost_with_cache}")
|
||||
|
||||
|
||||
def test_log_context_cost_calculation():
|
||||
"""
|
||||
Test that log context cost calculation works correctly with tiered pricing.
|
||||
|
||||
This test verifies that when using extended context (above 200k tokens),
|
||||
the log context costs are calculated using the appropriate tiered rates.
|
||||
"""
|
||||
from litellm import completion_cost
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
Message,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
# Create a mock response with extended context usage
|
||||
extended_context_response = ModelResponse(
|
||||
id="test-extended-context-response",
|
||||
created=1750733889,
|
||||
model="claude-4-sonnet-20250514",
|
||||
object="chat.completion",
|
||||
system_fingerprint=None,
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(
|
||||
content="This is a test response for extended context cost calculation.",
|
||||
role="assistant",
|
||||
tool_calls=None,
|
||||
function_call=None,
|
||||
),
|
||||
)
|
||||
],
|
||||
usage=Usage(
|
||||
total_tokens=350000, # Above 200k threshold
|
||||
prompt_tokens=301000, # Above 200k threshold
|
||||
completion_tokens=50000,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=300000,
|
||||
cached_tokens=0, # No cache hits
|
||||
audio_tokens=None,
|
||||
image_tokens=None,
|
||||
character_count=None,
|
||||
video_length_seconds=None,
|
||||
cache_creation_tokens=1000,
|
||||
),
|
||||
completion_tokens_details=None,
|
||||
_cache_creation_input_tokens=1000, # Some tokens added to cache
|
||||
),
|
||||
)
|
||||
|
||||
# Calculate the cost using the extended context model
|
||||
result = completion_cost(
|
||||
completion_response=extended_context_response,
|
||||
model="claude-4-sonnet-20250514",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
|
||||
# Debug: Print the actual result
|
||||
print(f"DEBUG: Actual cost result: ${result:.6f}")
|
||||
|
||||
# Get model info to understand the pricing
|
||||
from litellm import get_model_info
|
||||
|
||||
model_info = get_model_info(model="claude-4-sonnet-20250514", custom_llm_provider="anthropic")
|
||||
|
||||
# Calculate expected cost based on actual model pricing
|
||||
input_cost_per_token = model_info.get("input_cost_per_token", 0)
|
||||
output_cost_per_token = model_info.get("output_cost_per_token", 0)
|
||||
cache_creation_cost_per_token = model_info.get("cache_creation_input_token_cost", 0)
|
||||
|
||||
# Check if tiered pricing is applied
|
||||
input_cost_above_200k = model_info.get("input_cost_per_token_above_200k_tokens", input_cost_per_token)
|
||||
output_cost_above_200k = model_info.get("output_cost_per_token_above_200k_tokens", output_cost_per_token)
|
||||
cache_creation_above_200k = model_info.get(
|
||||
"cache_creation_input_token_cost_above_200k_tokens",
|
||||
cache_creation_cost_per_token,
|
||||
)
|
||||
|
||||
print(f"DEBUG: Base input cost per token: ${input_cost_per_token:.2e}")
|
||||
print(f"DEBUG: Base output cost per token: ${output_cost_per_token:.2e}")
|
||||
print(f"DEBUG: Base cache creation cost per token: ${cache_creation_cost_per_token:.2e}")
|
||||
|
||||
# Handle tiered pricing - if not available, use base pricing
|
||||
if input_cost_above_200k is not None:
|
||||
print(f"DEBUG: Tiered input cost per token (>200k): ${input_cost_above_200k:.2e}")
|
||||
else:
|
||||
print("DEBUG: No tiered input pricing available, using base pricing")
|
||||
input_cost_above_200k = input_cost_per_token
|
||||
|
||||
if output_cost_above_200k is not None:
|
||||
print(f"DEBUG: Tiered output cost per token (>200k): ${output_cost_above_200k:.2e}")
|
||||
else:
|
||||
print("DEBUG: No tiered output pricing available, using base pricing")
|
||||
output_cost_above_200k = output_cost_per_token
|
||||
|
||||
if cache_creation_above_200k is not None:
|
||||
print(f"DEBUG: Tiered cache creation cost per token (>200k): ${cache_creation_above_200k:.2e}")
|
||||
else:
|
||||
print("DEBUG: No tiered cache creation pricing available, using base pricing")
|
||||
cache_creation_above_200k = cache_creation_cost_per_token
|
||||
|
||||
# Since we're above 200k tokens, we should use tiered pricing if available
|
||||
expected_input_cost = 300000 * input_cost_above_200k
|
||||
expected_output_cost = 50000 * output_cost_above_200k
|
||||
expected_cache_cost = 1000 * cache_creation_above_200k
|
||||
expected_total = expected_input_cost + expected_output_cost + expected_cache_cost
|
||||
|
||||
print(f"DEBUG: Expected total: ${expected_total:.6f}")
|
||||
|
||||
# Allow for small floating point differences
|
||||
assert abs(result - expected_total) < 1e-6, f"Expected cost ${expected_total:.6f}, but got ${result:.6f}"
|
||||
|
||||
print(f"✓ Log context cost calculation with tiered pricing is correct: ${result:.6f}")
|
||||
print(f" - Input tokens (300k): ${expected_input_cost:.6f}")
|
||||
print(f" - Output tokens (50k): ${expected_output_cost:.6f}")
|
||||
print(f" - Cache creation (1k): ${expected_cache_cost:.6f}")
|
||||
print(f" - Total: ${result:.6f}")
|
||||
|
||||
|
||||
def test_gemini_25_explicit_caching_cost_direct_usage():
|
||||
|
|
@ -1814,56 +1571,6 @@ def test_cost_margin_with_discount(monkeypatch):
|
|||
print(f" - Expected: ${expected_cost:.6f}")
|
||||
|
||||
|
||||
def test_azure_image_generation_cost_calculator():
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.types.utils import (
|
||||
ImageObject,
|
||||
ImageResponse,
|
||||
ImageUsage,
|
||||
ImageUsageInputTokensDetails,
|
||||
)
|
||||
|
||||
response_cost_calculator_kwargs = {
|
||||
"response_object": ImageResponse(
|
||||
created=1761785270,
|
||||
background=None,
|
||||
data=[
|
||||
ImageObject(
|
||||
b64_json=None,
|
||||
revised_prompt="A futuristic, techno-inspired green duck wearing cool modern sunglasses. The duck has a sleek, metallic appearance with glowing neon green accents, standing on a high-tech urban background with holographic billboards and illuminated city lights in the distance. The duck's feathers have a glossy, high-tech sheen, resembling a robotic design but still maintaining its avian features. The scene has a vibrant, cyberpunk aesthetic with a neon color palette.",
|
||||
url="test-azure-blob-url-with-sas-token",
|
||||
)
|
||||
],
|
||||
output_format=None,
|
||||
quality="hd",
|
||||
size=None,
|
||||
usage=ImageUsage(
|
||||
input_tokens=0,
|
||||
input_tokens_details=ImageUsageInputTokensDetails(image_tokens=0, text_tokens=0),
|
||||
output_tokens=0,
|
||||
total_tokens=0,
|
||||
),
|
||||
),
|
||||
"model": "azure/dall-e-3",
|
||||
"cache_hit": False,
|
||||
"custom_llm_provider": "azure",
|
||||
"base_model": "azure/dall-e-3",
|
||||
"call_type": "aimage_generation",
|
||||
"optional_params": {},
|
||||
"custom_pricing": False,
|
||||
"prompt": "",
|
||||
"standard_built_in_tools_params": {
|
||||
"web_search_options": None,
|
||||
"file_search": None,
|
||||
},
|
||||
"router_model_id": "6738c432ffc9b733597c6b86613ca20dc5f49bde591fd3d03e7cd6aa25bb241e",
|
||||
"litellm_logging_obj": MagicMock(),
|
||||
"service_tier": None,
|
||||
}
|
||||
|
||||
cost = response_cost_calculator(**response_cost_calculator_kwargs)
|
||||
assert cost > 0.079
|
||||
|
||||
|
||||
def test_completion_cost_extracts_service_tier_from_response(_local_model_cost_map):
|
||||
|
|
@ -2616,87 +2323,6 @@ def test_gemini_without_cache_tokens_details():
|
|||
print("✅ Gemini without cacheTokensDetails works correctly")
|
||||
|
||||
|
||||
def test_gemini_implicit_caching_cost_calculation():
|
||||
"""
|
||||
Test for Issue #16341: Gemini implicit cached tokens not counted in spend log
|
||||
|
||||
When Gemini uses implicit caching, it returns cachedContentTokenCount but NOT
|
||||
cacheTokensDetails. In this case, we should subtract cachedContentTokenCount
|
||||
from text_tokens to correctly calculate costs.
|
||||
|
||||
See: https://github.com/BerriAI/litellm/issues/16341
|
||||
"""
|
||||
from litellm import completion_cost
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
# Simulate Gemini response with implicit caching (cachedContentTokenCount only)
|
||||
completion_response = {
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 10000,
|
||||
"candidatesTokenCount": 5,
|
||||
"totalTokenCount": 10005,
|
||||
"cachedContentTokenCount": 8000, # Implicit caching - no cacheTokensDetails
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 10000}],
|
||||
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 5}],
|
||||
}
|
||||
}
|
||||
|
||||
usage = VertexGeminiConfig._calculate_usage(completion_response)
|
||||
|
||||
# Verify parsing
|
||||
assert usage.cache_read_input_tokens == 8000, (
|
||||
f"cache_read_input_tokens should be 8000, got {usage.cache_read_input_tokens}"
|
||||
)
|
||||
assert usage.prompt_tokens_details.cached_tokens == 8000, (
|
||||
f"cached_tokens should be 8000, got {usage.prompt_tokens_details.cached_tokens}"
|
||||
)
|
||||
|
||||
# CRITICAL: text_tokens should be (10000 - 8000) = 2000, NOT 10000
|
||||
# This is the fix for issue #16341
|
||||
assert usage.prompt_tokens_details.text_tokens == 2000, (
|
||||
f"text_tokens should be 2000 (10000 - 8000), got {usage.prompt_tokens_details.text_tokens}"
|
||||
)
|
||||
|
||||
# Verify cost calculation uses cached token pricing
|
||||
response = ModelResponse(
|
||||
id="mock-id",
|
||||
model="gemini-2.0-flash",
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
message=Message(role="assistant", content="Hello!"),
|
||||
finish_reason="stop",
|
||||
)
|
||||
],
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=response,
|
||||
model="gemini-2.0-flash",
|
||||
custom_llm_provider="gemini",
|
||||
)
|
||||
|
||||
# Get model pricing for verification
|
||||
import litellm
|
||||
|
||||
model_info = litellm.get_model_info("gemini/gemini-2.0-flash")
|
||||
input_cost = model_info.get("input_cost_per_token", 0)
|
||||
cache_read_cost = model_info.get("cache_read_input_token_cost", input_cost)
|
||||
output_cost = model_info.get("output_cost_per_token", 0)
|
||||
|
||||
# Expected cost: (2000 * input) + (8000 * cache_read) + (5 * output)
|
||||
expected_cost = (2000 * input_cost) + (8000 * cache_read_cost) + (5 * output_cost)
|
||||
|
||||
assert abs(cost - expected_cost) < 1e-9, (
|
||||
f"Cost calculation is wrong. Got ${cost:.6f}, expected ${expected_cost:.6f}. "
|
||||
f"Cached tokens may not be using reduced pricing."
|
||||
)
|
||||
|
||||
print("✅ Issue #16341 fix verified: Gemini implicit caching cost calculated correctly")
|
||||
|
||||
|
||||
def test_additional_costs_only_for_azure_ai(_local_model_cost_map):
|
||||
|
|
@ -4847,3 +4473,27 @@ def test_cost_per_token_bedrock_qwen3_next_uses_regional_entry_not_us_rate(
|
|||
|
||||
assert prompt_usd == pytest.approx(prompt_tokens * regional["input_cost_per_token"])
|
||||
assert completion_usd == pytest.approx(completion_tokens * regional["output_cost_per_token"])
|
||||
|
||||
|
||||
def test_cost_per_token_bedrock_nemotron_super_3_uses_eu_west_2_entry_not_us_rate(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
|
||||
regional_key: Final = "bedrock/eu-west-2/nvidia.nemotron-super-3-120b"
|
||||
regional: Final = litellm.model_cost[regional_key]
|
||||
us: Final = litellm.model_cost["nvidia.nemotron-super-3-120b"]
|
||||
assert regional["input_cost_per_token"] != us["input_cost_per_token"]
|
||||
assert regional["output_cost_per_token"] != us["output_cost_per_token"]
|
||||
|
||||
prompt_tokens, completion_tokens = 1000, 500
|
||||
prompt_usd, completion_usd = cost_per_token(
|
||||
model=regional_key,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
assert prompt_usd == pytest.approx(prompt_tokens * regional["input_cost_per_token"])
|
||||
assert completion_usd == pytest.approx(completion_tokens * regional["output_cost_per_token"])
|
||||
|
|
|
|||
|
|
@ -157,21 +157,3 @@ def test_acount_tokens_no_api_key_falls_back(monkeypatch):
|
|||
assert result.tokenizer_type == "local_tokenizer"
|
||||
|
||||
|
||||
async def test_acount_tokens_local_fallback_counts_off_the_event_loop():
|
||||
from tests.large_text import text
|
||||
from tests.test_litellm.litellm_core_utils.event_loop_lag import (
|
||||
assert_loop_stayed_free,
|
||||
timed_with_loop_lags,
|
||||
warm_tokenizer,
|
||||
)
|
||||
|
||||
model = "together_ai/meta-llama/Llama-3-8b-chat-hf"
|
||||
warm_tokenizer(model)
|
||||
|
||||
result, took, lags = await timed_with_loop_lags(
|
||||
lambda: litellm.acount_tokens(model=model, messages=[{"role": "user", "content": text * 100}])
|
||||
)
|
||||
|
||||
assert result.tokenizer_type == "local_tokenizer"
|
||||
assert result.total_tokens > 100_000
|
||||
assert_loop_stayed_free(took, lags)
|
||||
|
|
|
|||
|
|
@ -90,27 +90,6 @@ class TestGPTImageCostCalculator:
|
|||
class TestGPTImageCostRouting:
|
||||
"""Test that gpt-image models are properly routed to the token-based calculator"""
|
||||
|
||||
def test_openai_dalle_routes_to_pixel_calculator(self):
|
||||
"""Test that OpenAI DALL-E still routes to pixel-based calculator"""
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import CostCalculatorUtils
|
||||
|
||||
image_response = ImageResponse(
|
||||
created=1234567890,
|
||||
data=[ImageObject(url="http://example.com/image.jpg")],
|
||||
)
|
||||
image_response.size = "1024x1024"
|
||||
image_response.quality = "standard"
|
||||
|
||||
cost = CostCalculatorUtils.route_image_generation_cost_calculator(
|
||||
model="dall-e-3",
|
||||
completion_response=image_response,
|
||||
custom_llm_provider="openai",
|
||||
size="1024x1024",
|
||||
quality="standard",
|
||||
n=1,
|
||||
)
|
||||
|
||||
assert cost >= 0
|
||||
|
||||
|
||||
class TestGPTImage15OutputImageTokens:
|
||||
|
|
|
|||
|
|
@ -850,36 +850,6 @@ async def test_arouter_async_get_healthy_deployments():
|
|||
assert result[0]["litellm_params"]["model"] == "gpt-3.5-turbo"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("litellm.amoderation")
|
||||
async def test_arouter_amoderation_with_credential_name(mock_amoderation):
|
||||
"""
|
||||
Test that router.amoderation passes litellm_credential_name to the underlying litellm.amoderation call
|
||||
"""
|
||||
mock_amoderation.return_value = AsyncMock()
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "text-moderation-stable",
|
||||
"litellm_params": {
|
||||
"model": "text-moderation-stable",
|
||||
"litellm_credential_name": "my-custom-auth",
|
||||
},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
await router.amoderation(input="I love everyone!", model="text-moderation-stable")
|
||||
|
||||
mock_amoderation.assert_called_once()
|
||||
call_kwargs = mock_amoderation.call_args[1] # Get the kwargs of the call
|
||||
print(
|
||||
"call kwargs for router.amoderation=",
|
||||
json.dumps(call_kwargs, indent=4, default=str),
|
||||
)
|
||||
assert call_kwargs["litellm_credential_name"] == "my-custom-auth"
|
||||
assert call_kwargs["model"] == "text-moderation-stable"
|
||||
|
||||
|
||||
def test_arouter_test_team_model():
|
||||
|
|
|
|||
|
|
@ -95,15 +95,6 @@ def _successor(info: dict[str, object]) -> str | None:
|
|||
return successor if isinstance(successor, str) else None
|
||||
|
||||
|
||||
def test_together_successor_metadata_points_at_known_models(cost_map: CostMap):
|
||||
successors = {
|
||||
model: successor
|
||||
for model, info in cost_map.items()
|
||||
if model.startswith("together_ai/") and (successor := _successor(info)) is not None
|
||||
}
|
||||
assert len(successors) >= 10
|
||||
for model, successor in successors.items():
|
||||
assert successor in cost_map, f"{model} names successor {successor} that is not in the map"
|
||||
|
||||
|
||||
def test_together_backup_cost_map_in_sync(cost_map: CostMap):
|
||||
|
|
|
|||
|
|
@ -1465,12 +1465,6 @@ class TestProxyFunctionCalling:
|
|||
("gemini/gemini-2.5-pro", "litellm_proxy/gemini/gemini-2.5-pro", True),
|
||||
("gemini/gemini-2.5-flash", "litellm_proxy/gemini/gemini-2.5-flash", True),
|
||||
# Groq models (mixed support)
|
||||
("groq/gemma-7b-it", "litellm_proxy/groq/gemma-7b-it", True),
|
||||
(
|
||||
"groq/llama-3.3-70b-versatile",
|
||||
"litellm_proxy/groq/llama-3.3-70b-versatile",
|
||||
True,
|
||||
),
|
||||
# Cohere models (generally don't support function calling)
|
||||
("command-nightly", "litellm_proxy/command-nightly", False),
|
||||
],
|
||||
|
|
|
|||
|
|
@ -46,34 +46,6 @@ class TestXAIResponsesAutoRouting:
|
|||
assert model_info.get("mode") != "responses"
|
||||
assert updated_model == model
|
||||
|
||||
def test_responses_api_bridge_check_with_tools(self):
|
||||
"""Test that with tools, xAI automatically routes to Responses API"""
|
||||
model = "grok-3"
|
||||
custom_llm_provider = "xai"
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get the weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
web_search_options = None
|
||||
|
||||
model_info, updated_model = responses_api_bridge_check(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
web_search_options=web_search_options,
|
||||
)
|
||||
|
||||
# Should auto-route to responses mode when tools are present
|
||||
assert model_info.get("mode") == "chat"
|
||||
assert updated_model == model
|
||||
|
||||
def test_responses_api_bridge_check_with_empty_tools(self):
|
||||
"""Test that with empty tools list, xAI does not route to Responses API"""
|
||||
|
|
@ -134,57 +106,8 @@ class TestXAIResponsesAutoRouting:
|
|||
assert model_info.get("mode") == "responses"
|
||||
assert updated_model == "grok-3" # prefix removed
|
||||
|
||||
def test_responses_api_bridge_check_with_code_interpreter_tool(self):
|
||||
"""Test auto-routing with code_interpreter tool"""
|
||||
model = "grok-3"
|
||||
custom_llm_provider = "xai"
|
||||
tools = [{"type": "code_interpreter"}]
|
||||
web_search_options = None
|
||||
|
||||
model_info, updated_model = responses_api_bridge_check(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
web_search_options=web_search_options,
|
||||
)
|
||||
# Should auto-route with code_interpreter tool
|
||||
assert model_info.get("mode") == "chat"
|
||||
assert updated_model == model
|
||||
|
||||
def test_responses_api_bridge_check_with_web_search_tool(self):
|
||||
"""Test auto-routing with web_search tool"""
|
||||
model = "grok-4"
|
||||
custom_llm_provider = "xai"
|
||||
tools = [
|
||||
{"type": "web_search", "filters": {"allowed_domains": ["wikipedia.org"]}}
|
||||
]
|
||||
web_search_options = None
|
||||
|
||||
model_info, updated_model = responses_api_bridge_check(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
web_search_options=web_search_options,
|
||||
)
|
||||
|
||||
# Should auto-route with web_search tool
|
||||
assert model_info.get("mode") == "chat"
|
||||
assert updated_model == model
|
||||
|
||||
def test_responses_api_bridge_check_with_x_search_tool(self):
|
||||
"""Test auto-routing with x_search tool"""
|
||||
model = "grok-4"
|
||||
custom_llm_provider = "xai"
|
||||
tools = [{"type": "x_search", "allowed_x_handles": ["@elonmusk"]}]
|
||||
web_search_options = None
|
||||
|
||||
model_info, updated_model = responses_api_bridge_check(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
web_search_options=web_search_options,
|
||||
)
|
||||
|
||||
# Should auto-route with x_search tool
|
||||
assert model_info.get("mode") == "chat"
|
||||
assert updated_model == model
|
||||
|
||||
def test_responses_api_bridge_check_with_web_search_options(self):
|
||||
"""Test auto-routing with web_search_options"""
|
||||
|
|
|
|||
|
|
@ -157,27 +157,6 @@ def test_foundry_gpt_6_astra_keeps_sampling_params_when_reasoning_effort_is_none
|
|||
assert optional_params == {"reasoning_effort": "none", "temperature": 0.2, "top_p": 0.9}
|
||||
|
||||
|
||||
def test_a_gpt_5_name_without_a_foundry_row_keeps_reading_its_own_entry(
|
||||
monkeypatch: pytest.MonkeyPatch, _local_model_cost_map
|
||||
):
|
||||
"""Most gpt-5-family names have no azure_ai/ row. Reading an azure_ai/ key for those finds
|
||||
nothing, and an openai.azure.com base sends the name down the azure provider, which has no key
|
||||
for it either, so every effort answer would silently fall back to false and take temperature,
|
||||
top_p and logprobs down with it."""
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", "https://example-resource.openai.azure.com")
|
||||
monkeypatch.setenv("AZURE_AI_API_KEY", "placeholder")
|
||||
|
||||
optional_params = litellm.utils.get_optional_params(
|
||||
model="gpt-5.1-chat-latest",
|
||||
custom_llm_provider="azure_ai",
|
||||
temperature=0.2,
|
||||
top_p=0.9,
|
||||
logprobs=True,
|
||||
)
|
||||
|
||||
assert optional_params["temperature"] == 0.2
|
||||
assert optional_params["top_p"] == 0.9
|
||||
assert optional_params["logprobs"] is True
|
||||
|
||||
|
||||
def test_azure_ai_grok_stop_parameter_handling():
|
||||
|
|
|
|||
|
|
@ -1695,21 +1695,6 @@ def test_nim_vllm_extras_translated_end_to_end_in_request_body():
|
|||
assert request_body["top_k"] == 40
|
||||
|
||||
|
||||
def test_in_schema_unsupported_params_still_raise():
|
||||
with pytest.raises(litellm.UnsupportedParamsError):
|
||||
litellm.get_optional_params(
|
||||
model="accounts/fireworks/models/llama-v3-70b-instruct",
|
||||
custom_llm_provider="fireworks_ai",
|
||||
drop_params=False,
|
||||
store=True,
|
||||
)
|
||||
optional_params = litellm.get_optional_params(
|
||||
model="accounts/fireworks/models/llama-v3-70b-instruct",
|
||||
custom_llm_provider="fireworks_ai",
|
||||
drop_params=True,
|
||||
store=True,
|
||||
)
|
||||
assert "store" not in optional_params
|
||||
|
||||
|
||||
def test_streaming_preserves_selected_model_for_private_accounting():
|
||||
|
|
|
|||
|
|
@ -715,7 +715,7 @@ class TestMoonshotReasoningEffort:
|
|||
def force_local_model_cost(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "model_cost", GetModelCostMap.load_local_model_cost_map())
|
||||
|
||||
@pytest.mark.parametrize("model", ["kimi-k3", "kimi-k2.5", "kimi-k2.6", "kimi-k2-thinking"])
|
||||
@pytest.mark.parametrize("model", ["kimi-k3", "kimi-k2.5", "kimi-k2.6"])
|
||||
def test_reasoning_model_supports_reasoning_effort(self, model):
|
||||
assert "reasoning_effort" in MoonshotChatConfig().get_supported_openai_params(model)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue