mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
Merge pull request #34670 from BerriAI/litellm_/stable-backport-1-93-x-d8c02a
chore(release): backport #33565, #33840, #33841, #34121, #33261, #34325 and #34577 to stable/1.93.x and cut 1.93.1
This commit is contained in:
commit
fc69949d74
16 changed files with 1078 additions and 137 deletions
|
|
@ -54,7 +54,6 @@ ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
|
|||
UV_LINK_MODE=copy \
|
||||
PATH="/app/.venv/bin:${PATH}" \
|
||||
LITELLM_NON_ROOT=true \
|
||||
PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
|
||||
XDG_CACHE_HOME=/app/.cache
|
||||
|
||||
# Copy dependency metadata first for layer caching
|
||||
|
|
@ -106,7 +105,9 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
|
|||
--python python3; \
|
||||
fi
|
||||
|
||||
RUN prisma generate --schema=./schema.prisma
|
||||
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
npm_config_cache=/root/.npm \
|
||||
prisma generate --schema=./schema.prisma
|
||||
|
||||
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
|
||||
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
|
||||
|
|
@ -127,8 +128,6 @@ RUN for i in 1 2 3; do \
|
|||
# the rest of the builder's /app is source and build metadata that must not
|
||||
# ship (manifest-scanning tools attribute everything in it to this image).
|
||||
# entrypoint.sh invokes litellm/proxy/prisma_migration.py by source path.
|
||||
# Prisma caches live under /app/.cache here (XDG_CACHE_HOME /
|
||||
# PRISMA_BINARY_CACHE_DIR) so the runtime prisma generate finds them.
|
||||
COPY --from=builder /app/.venv /app/.venv
|
||||
COPY --from=builder /app/docker /app/docker
|
||||
COPY --from=builder /app/schema.prisma /app/schema.prisma
|
||||
|
|
@ -138,21 +137,35 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr
|
|||
# enterprise.enterprise_hooks from it)
|
||||
COPY --from=builder /app/enterprise /app/enterprise
|
||||
COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras
|
||||
COPY --from=builder /app/.cache /app/.cache
|
||||
# Prisma CLI + engines are baked under /opt/prisma, a fixed path every runtime
|
||||
# uid can read and that no cache volume mount shadows (unlike /app/.cache or
|
||||
# $HOME/.cache under readOnlyRootFilesystem + emptyDir or arbitrary-uid setups).
|
||||
# PRISMA_CLI_QUERY_ENGINE_TYPE=binary makes the CLI use the baked binary query
|
||||
# engine directly, so `prisma migrate deploy` on a fresh database needs no npm
|
||||
# and no network access; without it the CLI looks for the library engine, which
|
||||
# prisma stopped baking, and falls back to a download that fails offline or as a
|
||||
# non-writable uid (#33650, #24554).
|
||||
COPY --from=builder /opt/prisma /opt/prisma
|
||||
COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui
|
||||
COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets
|
||||
|
||||
# XDG_CACHE_HOME is intentionally left unset so it falls back to $HOME/.cache
|
||||
# (/app/.cache, writable by the runtime uid). The prisma bake at the read-only
|
||||
# /opt/prisma is anchored by PRISMA_BINARY_CACHE_DIR / PRISMA_CLI_PATH, so
|
||||
# nothing needs XDG to point there; pointing it at the read-only bake would
|
||||
# deny any XDG-aware library that writes a cache at runtime.
|
||||
ENV PATH="/app/.venv/bin:${PATH}" \
|
||||
PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \
|
||||
PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \
|
||||
PRISMA_CLI_PATH=/opt/prisma/binaries/node_modules/.bin/prisma \
|
||||
PRISMA_CLI_QUERY_ENGINE_TYPE=binary \
|
||||
HOME=/app \
|
||||
LITELLM_NON_ROOT=true \
|
||||
XDG_CACHE_HOME=/app/.cache \
|
||||
PRISMA_SKIP_POSTINSTALL_GENERATE=1 \
|
||||
PRISMA_HIDE_UPDATE_MESSAGE=1 \
|
||||
PRISMA_ENGINES_CHECKSUM_IGNORE_MISSING=1 \
|
||||
PRISMA_OFFLINE_MODE=true
|
||||
|
||||
RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \
|
||||
RUN mkdir -p /nonexistent /app/.cache /var/lib/litellm/assets /var/lib/litellm/ui && \
|
||||
chown -R nobody:nogroup /app /var/lib/litellm/ui /var/lib/litellm/assets /nonexistent && \
|
||||
PRISMA_PATH=$(python -c "import os, prisma; print(os.path.dirname(prisma.__file__))") && \
|
||||
chown -R nobody:nogroup "$PRISMA_PATH" && \
|
||||
|
|
@ -165,12 +178,14 @@ RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \
|
|||
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g=u "$LITELLM_PROXY_EXTRAS_PATH" || true && \
|
||||
chmod -R g+w "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets && \
|
||||
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g+w "$LITELLM_PROXY_EXTRAS_PATH" || true && \
|
||||
chmod -R g+rX "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets /app/.cache
|
||||
chmod -R g+rX "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets && \
|
||||
chmod -R a+rX /opt/prisma && \
|
||||
test -x /opt/prisma/binaries/node_modules/.bin/prisma && \
|
||||
test -f /opt/prisma/binaries/node_modules/prisma/build/index.js && \
|
||||
ls /opt/prisma/binaries/node_modules/@prisma/engines/query-engine-* >/dev/null 2>&1
|
||||
|
||||
USER 65534
|
||||
|
||||
RUN prisma generate --schema=./schema.prisma
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
ENTRYPOINT ["/app/docker/prod_entrypoint.sh"]
|
||||
|
|
|
|||
|
|
@ -103,6 +103,18 @@ class DualCache(BaseCache):
|
|||
if default_redis_ttl is not None:
|
||||
self.default_redis_ttl = default_redis_ttl
|
||||
|
||||
def _backfill_kwargs(self, kwargs: "dict[str, object]") -> "dict[str, object]":
|
||||
"""
|
||||
Kwargs for writing a Redis read result into the in-memory tier.
|
||||
|
||||
Applies ``default_in_memory_ttl`` exactly like the write paths do;
|
||||
without it, backfilled entries fall to ``InMemoryCache``'s own default
|
||||
TTL and can outlive the TTL this cache was configured with.
|
||||
"""
|
||||
if "ttl" not in kwargs and self.default_in_memory_ttl is not None:
|
||||
return {**kwargs, "ttl": self.default_in_memory_ttl}
|
||||
return kwargs
|
||||
|
||||
def set_cache(self, key, value, local_only: bool = False, **kwargs):
|
||||
# Update both Redis and in-memory cache
|
||||
try:
|
||||
|
|
@ -160,7 +172,7 @@ class DualCache(BaseCache):
|
|||
|
||||
if redis_result is not None:
|
||||
# Update in-memory cache with the value from Redis
|
||||
self.in_memory_cache.set_cache(key, redis_result, **kwargs)
|
||||
self.in_memory_cache.set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
|
||||
|
||||
result = redis_result
|
||||
|
||||
|
|
@ -226,7 +238,7 @@ class DualCache(BaseCache):
|
|||
|
||||
if redis_result is not None:
|
||||
# Update in-memory cache with the value from Redis
|
||||
await self.in_memory_cache.async_set_cache(key, redis_result, **kwargs)
|
||||
await self.in_memory_cache.async_set_cache(key, redis_result, **self._backfill_kwargs(kwargs))
|
||||
|
||||
result = redis_result
|
||||
|
||||
|
|
@ -318,7 +330,7 @@ class DualCache(BaseCache):
|
|||
result[key_to_index[key]] = value
|
||||
|
||||
if value is not None and self.in_memory_cache is not None:
|
||||
await self.in_memory_cache.async_set_cache(key, value, **kwargs)
|
||||
await self.in_memory_cache.async_set_cache(key, value, **self._backfill_kwargs(kwargs))
|
||||
|
||||
return result
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -1010,7 +1010,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None:
|
|||
return
|
||||
if getattr(request.state, "parent_otel_span", None) is not None:
|
||||
return
|
||||
start_time = datetime.now()
|
||||
start_time = datetime.now(timezone.utc)
|
||||
try:
|
||||
request.state.litellm_received_at = start_time
|
||||
except Exception:
|
||||
|
|
@ -1060,7 +1060,7 @@ async def _user_api_key_auth_builder(
|
|||
# Prefer the receive-instant stamped by the early helper in
|
||||
# user_api_key_auth (before body parse) — overwriting it would shorten
|
||||
# the preprocessing-duration measurement by the body-parse window.
|
||||
start_time = getattr(request.state, "litellm_received_at", None) or datetime.now()
|
||||
start_time = getattr(request.state, "litellm_received_at", None) or datetime.now(timezone.utc)
|
||||
try:
|
||||
request.state.litellm_received_at = start_time
|
||||
except Exception:
|
||||
|
|
@ -1670,10 +1670,9 @@ async def _user_api_key_auth_builder(
|
|||
valid_token.end_user_tpm_limit = end_user_params.get("end_user_tpm_limit")
|
||||
valid_token.end_user_rpm_limit = end_user_params.get("end_user_rpm_limit")
|
||||
valid_token.allowed_model_region = end_user_params.get("allowed_model_region")
|
||||
# update key budget with temp budget increase
|
||||
valid_token = _update_key_budget_with_temp_budget_increase(
|
||||
valid_token
|
||||
) # updating it here, allows all downstream reporting / checks to use the updated budget
|
||||
|
||||
if valid_token is not None:
|
||||
valid_token = _update_key_budget_with_temp_budget_increase(valid_token)
|
||||
|
||||
user_obj: Optional[LiteLLM_UserTable] = None
|
||||
valid_token_dict: dict = {}
|
||||
|
|
@ -1992,16 +1991,6 @@ async def _user_api_key_auth_builder(
|
|||
raise HTTPException(401, detail="Invalid API key, no token associated")
|
||||
api_key = valid_token.token
|
||||
|
||||
# Add hashed token to cache
|
||||
asyncio.create_task(
|
||||
_cache_key_object(
|
||||
hashed_token=api_key,
|
||||
user_api_key_obj=valid_token,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
|
||||
valid_token_dict = valid_token.model_dump(exclude_none=True)
|
||||
valid_token_dict.pop("token", None)
|
||||
# budget_throttle_pct is excluded from model_dump (it must not leak
|
||||
|
|
@ -2311,6 +2300,9 @@ async def _run_centralized_common_checks(
|
|||
None if isinstance(global_spend_result, BaseException) else global_spend_result
|
||||
)
|
||||
|
||||
if user_api_key_auth_obj.org_id is None and team_object is not None and team_object.organization_id is not None:
|
||||
user_api_key_auth_obj.org_id = team_object.organization_id
|
||||
|
||||
# common_checks identifies admin via user_object, not the token
|
||||
# (non_proxy_admin_allowed_routes_check). JWT admin shortcut and
|
||||
# master_key tokens get admin from the token; the DB row for the
|
||||
|
|
@ -2612,7 +2604,7 @@ async def _return_user_api_key_auth_obj(
|
|||
start_time: datetime,
|
||||
user_role: Optional[LitellmUserRoles] = None,
|
||||
) -> UserAPIKeyAuth:
|
||||
end_time = datetime.now()
|
||||
end_time = datetime.now(timezone.utc)
|
||||
|
||||
asyncio.create_task(
|
||||
user_api_key_service_logger_obj.async_service_success_hook(
|
||||
|
|
@ -2689,7 +2681,9 @@ def _get_temp_budget_increase(valid_token: UserAPIKeyAuth):
|
|||
valid_token_metadata = valid_token.metadata
|
||||
if "temp_budget_increase" in valid_token_metadata and "temp_budget_expiry" in valid_token_metadata:
|
||||
expiry = datetime.fromisoformat(valid_token_metadata["temp_budget_expiry"])
|
||||
if expiry > datetime.now():
|
||||
if expiry.tzinfo is None:
|
||||
expiry = expiry.replace(tzinfo=timezone.utc)
|
||||
if expiry > datetime.now(timezone.utc):
|
||||
return valid_token_metadata["temp_budget_increase"]
|
||||
return None
|
||||
|
||||
|
|
@ -2699,9 +2693,10 @@ def _update_key_budget_with_temp_budget_increase(
|
|||
) -> UserAPIKeyAuth:
|
||||
if valid_token.max_budget is None:
|
||||
return valid_token
|
||||
temp_budget_increase = _get_temp_budget_increase(valid_token) or 0.0
|
||||
valid_token.max_budget = valid_token.max_budget + temp_budget_increase
|
||||
return valid_token
|
||||
temp_budget_increase = _get_temp_budget_increase(valid_token)
|
||||
if not temp_budget_increase:
|
||||
return valid_token
|
||||
return valid_token.model_copy(update={"max_budget": valid_token.max_budget + temp_budget_increase})
|
||||
|
||||
|
||||
async def _lookup_end_user_and_apply_budget(
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import asyncio
|
|||
import base64
|
||||
import hashlib
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
|
|
@ -245,18 +246,28 @@ def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dic
|
|||
raise HTTPException(status_code=400, detail="Invalid CLI login session")
|
||||
|
||||
cache_key = _get_cli_sso_flow_cache_key(cast(str, login_id))
|
||||
flow = cache.get_cache(key=cache_key)
|
||||
redis_cache = cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
flow = redis_cache.get_cache(key=cache_key)
|
||||
else:
|
||||
flow = cache.get_cache(key=cache_key)
|
||||
if isinstance(flow, str):
|
||||
try:
|
||||
flow = json.loads(flow)
|
||||
except ValueError:
|
||||
flow = None
|
||||
if not isinstance(flow, dict) or "poll_secret_hash" not in flow:
|
||||
raise HTTPException(status_code=400, detail="Invalid CLI login session")
|
||||
return flow
|
||||
|
||||
|
||||
def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None:
|
||||
cache.set_cache(
|
||||
key=_get_cli_sso_flow_cache_key(login_id),
|
||||
value=flow,
|
||||
ttl=CLI_SSO_SESSION_TTL_SECONDS,
|
||||
)
|
||||
cache_key = _get_cli_sso_flow_cache_key(login_id)
|
||||
redis_cache = cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
redis_cache.set_cache(key=cache_key, value=json.dumps(flow), ttl=CLI_SSO_SESSION_TTL_SECONDS)
|
||||
else:
|
||||
cache.set_cache(key=cache_key, value=flow, ttl=CLI_SSO_SESSION_TTL_SECONDS)
|
||||
|
||||
|
||||
def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool:
|
||||
|
|
@ -567,11 +578,11 @@ def _render_cli_sso_verification_page(
|
|||
|
||||
@router.post("/sso/cli/start", tags=["experimental"], include_in_schema=False)
|
||||
async def cli_sso_start(request: Request):
|
||||
from litellm.proxy.proxy_server import general_settings, user_api_key_cache
|
||||
from litellm.proxy.proxy_server import cli_sso_session_cache, general_settings
|
||||
|
||||
_check_cli_sso_start_rate_limit(
|
||||
request=request,
|
||||
cache=user_api_key_cache,
|
||||
cache=cli_sso_session_cache,
|
||||
use_x_forwarded_for=bool((general_settings or {}).get("use_x_forwarded_for", False)),
|
||||
)
|
||||
|
||||
|
|
@ -586,7 +597,7 @@ async def cli_sso_start(request: Request):
|
|||
"user_code_verified": False,
|
||||
"session_data": None,
|
||||
}
|
||||
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
|
||||
_set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow)
|
||||
|
||||
verification_uri_complete: str | None = (
|
||||
(
|
||||
|
|
@ -618,9 +629,9 @@ async def cli_sso_complete(request: Request, login_id: str):
|
|||
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
|
||||
render_cli_sso_success_page,
|
||||
)
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
from litellm.proxy.proxy_server import cli_sso_session_cache
|
||||
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=user_api_key_cache)
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cli_sso_session_cache)
|
||||
if not flow.get("sso_complete") or not flow.get("session_data"):
|
||||
raise HTTPException(status_code=400, detail="CLI login is not ready")
|
||||
|
||||
|
|
@ -644,7 +655,7 @@ async def cli_sso_complete(request: Request, login_id: str):
|
|||
raise HTTPException(status_code=400, detail="Invalid verification code")
|
||||
|
||||
flow["user_code_verified"] = True
|
||||
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
|
||||
_set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow)
|
||||
|
||||
html_content = render_cli_sso_success_page()
|
||||
return HTMLResponse(content=html_content, status_code=200)
|
||||
|
|
@ -835,10 +846,10 @@ async def google_login(
|
|||
Example:
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
cli_sso_session_cache,
|
||||
general_settings,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
user_custom_ui_sso_sign_in_handler,
|
||||
)
|
||||
|
||||
|
|
@ -886,7 +897,7 @@ async def google_login(
|
|||
)
|
||||
|
||||
if source == LITELLM_CLI_SOURCE_IDENTIFIER:
|
||||
_get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
|
||||
_get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache)
|
||||
|
||||
# Store CLI login handle in state for OAuth flow
|
||||
cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state(
|
||||
|
|
@ -1920,6 +1931,7 @@ async def _complete_cli_sso_callback_session(
|
|||
user_defined_values: Optional[SSOUserDefinedValues],
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
cli_sso_session_cache: DualCache,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
prefill_user_code: str | None = None,
|
||||
):
|
||||
|
|
@ -1966,7 +1978,7 @@ async def _complete_cli_sso_callback_session(
|
|||
flow["sso_complete"] = True
|
||||
browser_complete_token = secrets.token_urlsafe(32)
|
||||
flow["browser_complete_token_hash"] = _hash_cli_sso_secret(browser_complete_token)
|
||||
_set_cli_sso_flow(login_id=key, cache=user_api_key_cache, flow=flow)
|
||||
_set_cli_sso_flow(login_id=key, cache=cli_sso_session_cache, flow=flow)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}"
|
||||
|
|
@ -1996,13 +2008,14 @@ async def cli_sso_callback(
|
|||
verbose_proxy_logger.info("CLI SSO callback")
|
||||
|
||||
from litellm.proxy.proxy_server import (
|
||||
cli_sso_session_cache,
|
||||
general_settings,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
|
@ -2042,6 +2055,7 @@ async def cli_sso_callback(
|
|||
user_defined_values=user_defined_values,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
cli_sso_session_cache=cli_sso_session_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
prefill_user_code=prefill_user_code,
|
||||
)
|
||||
|
|
@ -2076,10 +2090,14 @@ async def cli_poll_key(
|
|||
get_team_object,
|
||||
get_user_object,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
from litellm.proxy.proxy_server import (
|
||||
cli_sso_session_cache,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
try:
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache)
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=cli_sso_session_cache)
|
||||
if not _verify_cli_sso_poll_secret(flow=flow, poll_secret=x_litellm_cli_poll_secret):
|
||||
raise HTTPException(status_code=403, detail="Invalid CLI polling secret")
|
||||
|
||||
|
|
@ -2186,7 +2204,7 @@ async def cli_poll_key(
|
|||
)
|
||||
|
||||
# Delete cache entry (single-use)
|
||||
user_api_key_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id))
|
||||
cli_sso_session_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id))
|
||||
|
||||
verbose_proxy_logger.info(f"CLI JWT generated for user: {user_id}, team: {team_id}")
|
||||
poll_response = {
|
||||
|
|
|
|||
|
|
@ -224,6 +224,7 @@ from litellm.constants import (
|
|||
APSCHEDULER_MAX_INSTANCES,
|
||||
APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
APSCHEDULER_REPLACE_EXISTING,
|
||||
CLI_SSO_SESSION_TTL_SECONDS,
|
||||
DAYS_IN_A_MONTH,
|
||||
DEFAULT_HEALTH_CHECK_INTERVAL,
|
||||
DEFAULT_MODEL_CREATED_AT_TIME,
|
||||
|
|
@ -1909,6 +1910,7 @@ user_api_key_cache: UserApiKeyCache = UserApiKeyCache(
|
|||
default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
|
||||
)
|
||||
spend_counter_cache = DualCache(default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value)
|
||||
cli_sso_session_cache = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_TTL_SECONDS)
|
||||
model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=user_api_key_cache)
|
||||
litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter)
|
||||
redis_usage_cache: Optional[RedisCache] = None # redis cache used for tracking spend, tpm/rpm limits
|
||||
|
|
@ -2768,21 +2770,6 @@ async def update_cache(
|
|||
)
|
||||
# set cooldown on alert
|
||||
|
||||
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_spend", None) is not None:
|
||||
existing_team_spend = existing_spend_obj.team_spend or 0
|
||||
# Calculate the new cost by adding the existing cost and response_cost
|
||||
existing_spend_obj.team_spend = existing_team_spend + response_cost
|
||||
|
||||
if existing_spend_obj is not None and getattr(existing_spend_obj, "team_member_spend", None) is not None:
|
||||
existing_team_member_spend = existing_spend_obj.team_member_spend or 0
|
||||
# Calculate the new cost by adding the existing cost and response_cost
|
||||
existing_spend_obj.team_member_spend = existing_team_member_spend + response_cost
|
||||
|
||||
# Existing spend_obj is mutated; UserApiKeyCache.async_set_cache_pipeline turns
|
||||
# BaseModel values into dicts for Redis (same Codec path as async_set_cache).
|
||||
existing_spend_obj.spend = new_spend
|
||||
values_to_update_in_cache.append((hashed_token, existing_spend_obj))
|
||||
|
||||
### UPDATE USER SPEND ###
|
||||
async def _update_user_cache():
|
||||
## UPDATE CACHE FOR USER ID + GLOBAL PROXY
|
||||
|
|
@ -2986,13 +2973,27 @@ async def update_cache(
|
|||
if tags is not None:
|
||||
await _update_tag_cache()
|
||||
|
||||
asyncio.create_task(
|
||||
user_api_key_cache.async_set_cache_pipeline(
|
||||
cache_list=values_to_update_in_cache,
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
global_proxy_spend_key = "{}:spend".format(litellm_proxy_admin_name)
|
||||
local_object_updates = tuple((k, v) for k, v in values_to_update_in_cache if k != global_proxy_spend_key)
|
||||
shared_scalar_updates = tuple((k, v) for k, v in values_to_update_in_cache if k == global_proxy_spend_key)
|
||||
|
||||
if local_object_updates:
|
||||
asyncio.create_task(
|
||||
user_api_key_cache.async_set_cache_pipeline(
|
||||
cache_list=list(local_object_updates),
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
local_only=True,
|
||||
)
|
||||
)
|
||||
if shared_scalar_updates:
|
||||
asyncio.create_task(
|
||||
user_api_key_cache.async_set_cache_pipeline(
|
||||
cache_list=list(shared_scalar_updates),
|
||||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
litellm_parent_otel_span=parent_otel_span,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def run_ollama_serve():
|
||||
|
|
@ -3633,13 +3634,22 @@ def _build_redis_usage_cache_from_environment() -> RedisCache | None:
|
|||
def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache: bool) -> None:
|
||||
"""
|
||||
Wires an established coordination Redis into the proxy-level caches that
|
||||
consume it directly: the spend counter cache, the cluster-wide config
|
||||
cache, and (only when opted in) the virtual-key auth cache.
|
||||
consume it directly: the spend counter cache, the CLI SSO login-session
|
||||
cache, the cluster-wide config cache, and (only when opted in) the
|
||||
virtual-key auth cache.
|
||||
|
||||
The CLI SSO login-session cache is always backed by Redis when available so
|
||||
that the browser SSO flow behind `lite login` survives landing on different
|
||||
workers; it must not be gated behind enable_redis_auth_cache.
|
||||
"""
|
||||
spend_counter_cache.attach_redis_cache(
|
||||
redis_cache,
|
||||
default_redis_ttl=litellm.default_redis_ttl,
|
||||
)
|
||||
cli_sso_session_cache.attach_redis_cache(
|
||||
redis_cache,
|
||||
default_redis_ttl=CLI_SSO_SESSION_TTL_SECONDS,
|
||||
)
|
||||
if enable_redis_auth_cache is True:
|
||||
user_api_key_cache.attach_redis_cache(
|
||||
redis_cache,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm"
|
||||
version = "1.93.0"
|
||||
version = "1.93.1"
|
||||
description = "Library to easily interface with LLM API providers"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10, <3.15"
|
||||
|
|
@ -286,7 +286,7 @@ members = ["enterprise", "litellm-proxy-extras"]
|
|||
profile = "black"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.93.0"
|
||||
version = "1.93.1"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@
|
|||
"limit": 33
|
||||
},
|
||||
"DTZ005": {
|
||||
"limit": 244
|
||||
"limit": 241
|
||||
},
|
||||
"DTZ006": {
|
||||
"limit": 13
|
||||
|
|
|
|||
143
tests/proxy_migration_tests/test_offline_image_migration.py
Normal file
143
tests/proxy_migration_tests/test_offline_image_migration.py
Normal file
|
|
@ -0,0 +1,143 @@
|
|||
"""Image-level regression net for the prisma bake in the shipped runtime image.
|
||||
|
||||
Boots a built image's migration entrypoint the way an OpenShift / air-gapped
|
||||
deployment does (an internal-only network with no egress, an arbitrary non-root
|
||||
uid in GID 0) against a brand-new Postgres, and asserts the schema was created.
|
||||
|
||||
This catches the whole failure class, not one symptom: a bake that only works
|
||||
under `docker run` as the default uid with network still passes every existing
|
||||
check, because the migration entrypoint exits 0 even when it applied nothing.
|
||||
Asserting the table count is what turns that silent success into a hard fail.
|
||||
|
||||
Gated on LITELLM_IMAGE (the tag of the image to exercise) so it is skipped in
|
||||
the normal unit-test run and exercised only where an image has been built (the
|
||||
image-scan workflow). Requires a working docker CLI.
|
||||
"""
|
||||
|
||||
import shutil
|
||||
import subprocess
|
||||
import uuid
|
||||
|
||||
import os
|
||||
import pytest
|
||||
|
||||
IMAGE = os.getenv("LITELLM_IMAGE")
|
||||
POSTGRES_IMAGE = os.getenv("LITELLM_TEST_POSTGRES_IMAGE", "postgres:16-alpine")
|
||||
MIN_TABLES = int(os.getenv("LITELLM_TEST_MIN_TABLES", "20"))
|
||||
NON_ROOT_UID = "12345:0" # arbitrary uid in GID 0, as OpenShift restricted-v2 assigns
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(IMAGE is None, reason="requires a built image (set LITELLM_IMAGE)"),
|
||||
pytest.mark.skipif(shutil.which("docker") is None, reason="requires the docker CLI"),
|
||||
]
|
||||
|
||||
|
||||
def _docker(*args: str, check: bool = True) -> subprocess.CompletedProcess:
|
||||
return subprocess.run(
|
||||
["docker", *args], capture_output=True, text=True, check=check
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def offline_postgres():
|
||||
"""A fresh Postgres reachable only over an internal-only (no egress) network.
|
||||
|
||||
Yields (network_name, postgres_host). Both are torn down afterwards.
|
||||
"""
|
||||
run_id = f"offlinemig-{uuid.uuid4().hex[:8]}"
|
||||
network = f"{run_id}-net"
|
||||
pg = f"{run_id}-pg"
|
||||
|
||||
# Pull Postgres while egress still exists; the internal network below has none.
|
||||
_docker("pull", "--quiet", POSTGRES_IMAGE)
|
||||
# --internal => containers on this network cannot reach the internet, so a
|
||||
# prisma engine download (binaries.prisma.sh / npm) fails instead of masking
|
||||
# a non-self-contained bake.
|
||||
_docker("network", "create", "--internal", network)
|
||||
try:
|
||||
_docker(
|
||||
"run", "-d", "--name", pg, "--network", network,
|
||||
"-e", "POSTGRES_PASSWORD=pw", "-e", "POSTGRES_DB=litellm",
|
||||
POSTGRES_IMAGE,
|
||||
)
|
||||
_wait_until_ready(pg)
|
||||
yield network, pg
|
||||
finally:
|
||||
_docker("rm", "-f", pg, check=False)
|
||||
_docker("network", "rm", network, check=False)
|
||||
|
||||
|
||||
def _wait_until_ready(pg: str, attempts: int = 60) -> None:
|
||||
for _ in range(attempts):
|
||||
running = _docker(
|
||||
"ps", "--filter", f"name={pg}", "--filter", "status=running",
|
||||
"--format", "{{.Names}}", check=False,
|
||||
).stdout
|
||||
if pg not in running:
|
||||
logs = _docker("logs", pg, check=False).stdout + _docker("logs", pg, check=False).stderr
|
||||
pytest.fail(f"postgres container is not running:\n{logs}")
|
||||
ready = _docker(
|
||||
"exec", pg, "pg_isready", "-U", "postgres", "-d", "litellm", check=False
|
||||
)
|
||||
if ready.returncode == 0:
|
||||
return
|
||||
subprocess.run(["sleep", "1"])
|
||||
pytest.fail(f"postgres never became ready after {attempts}s")
|
||||
|
||||
|
||||
def _table_count(pg: str) -> int:
|
||||
result = _docker(
|
||||
"exec", pg, "psql", "-U", "postgres", "-d", "litellm", "-tAc",
|
||||
"SELECT count(*) FROM information_schema.tables WHERE table_schema='public';",
|
||||
)
|
||||
return int(result.stdout.strip() or "0")
|
||||
|
||||
|
||||
def test_migration_offline_as_non_root_uid(offline_postgres):
|
||||
"""The migration entrypoint creates the full schema offline as an arbitrary uid.
|
||||
|
||||
Reproduces the OpenShift / air-gapped failure: on the pre-fix image the
|
||||
migration exits 0 having created 0 tables (every DB endpoint then 500s on
|
||||
missing columns); a self-contained bake creates the full schema.
|
||||
"""
|
||||
network, pg = offline_postgres
|
||||
assert IMAGE is not None
|
||||
|
||||
migrate = _docker(
|
||||
"run", "--rm", "--network", network, "--user", NON_ROOT_UID,
|
||||
"-e", f"DATABASE_URL=postgresql://postgres:pw@{pg}:5432/litellm",
|
||||
"-e", "LITELLM_MASTER_KEY=sk-offline-migration-test",
|
||||
"-e", "DISABLE_SCHEMA_UPDATE=false",
|
||||
"-w", "/app", "--entrypoint", "python",
|
||||
IMAGE, "litellm/proxy/prisma_migration.py",
|
||||
check=False,
|
||||
)
|
||||
tables = _table_count(pg)
|
||||
|
||||
assert migrate.returncode == 0, (
|
||||
f"migration entrypoint exited {migrate.returncode} offline as uid {NON_ROOT_UID}\n"
|
||||
f"stdout:\n{migrate.stdout}\nstderr:\n{migrate.stderr}"
|
||||
)
|
||||
assert tables >= MIN_TABLES, (
|
||||
f"only {tables} tables created (need >= {MIN_TABLES}) offline as uid {NON_ROOT_UID}. "
|
||||
"The prisma bake is not self-contained: it needs a runtime download or a "
|
||||
"writable HOME/cache, so OpenShift and air-gapped deployments start on an "
|
||||
f"empty database.\nstdout:\n{migrate.stdout}\nstderr:\n{migrate.stderr}"
|
||||
)
|
||||
|
||||
|
||||
def test_runtime_cache_env_not_read_only():
|
||||
"""No runtime cache env var may point at the world-read-only /opt/prisma bake.
|
||||
|
||||
/opt/prisma is baked `a+rX` (no write). Pointing XDG_CACHE_HOME (or any cache
|
||||
var an XDG-aware library honours) there would deny writes for every uid, so
|
||||
guard against a future edit reintroducing that.
|
||||
"""
|
||||
assert IMAGE is not None
|
||||
env = _docker("run", "--rm", "--entrypoint", "env", IMAGE).stdout
|
||||
offenders = [
|
||||
line for line in env.splitlines()
|
||||
if line.startswith(("XDG_CACHE_HOME=", "XDG_DATA_HOME=", "HOME="))
|
||||
and line.split("=", 1)[1].startswith("/opt/prisma")
|
||||
]
|
||||
assert not offenders, f"cache/home env points at the read-only bake: {offenders}"
|
||||
|
|
@ -1732,6 +1732,35 @@ def test_get_temp_budget_increase():
|
|||
assert _get_temp_budget_increase(valid_token) == 100
|
||||
|
||||
|
||||
def test_get_temp_budget_increase_tz_aware_expiry():
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import _get_temp_budget_increase
|
||||
|
||||
future_expiry = (datetime.now(timezone.utc) + timedelta(days=1)).isoformat()
|
||||
valid_token = UserAPIKeyAuth(
|
||||
max_budget=100,
|
||||
spend=0,
|
||||
metadata={
|
||||
"temp_budget_increase": 100,
|
||||
"temp_budget_expiry": future_expiry,
|
||||
},
|
||||
)
|
||||
assert _get_temp_budget_increase(valid_token) == 100
|
||||
|
||||
past_expiry = (datetime.now(timezone.utc) - timedelta(days=1)).isoformat()
|
||||
expired_token = UserAPIKeyAuth(
|
||||
max_budget=100,
|
||||
spend=0,
|
||||
metadata={
|
||||
"temp_budget_increase": 100,
|
||||
"temp_budget_expiry": past_expiry,
|
||||
},
|
||||
)
|
||||
assert _get_temp_budget_increase(expired_token) is None
|
||||
|
||||
|
||||
def test_update_key_budget_with_temp_budget_increase():
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
|
|
@ -1751,7 +1780,10 @@ def test_update_key_budget_with_temp_budget_increase():
|
|||
"temp_budget_expiry": expiry_in_isoformat,
|
||||
},
|
||||
)
|
||||
assert _update_key_budget_with_temp_budget_increase(valid_token).max_budget == 200
|
||||
result = _update_key_budget_with_temp_budget_increase(valid_token)
|
||||
assert result.max_budget == 200
|
||||
assert result is not valid_token
|
||||
assert valid_token.max_budget == 100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -88,6 +88,60 @@ async def test_dual_cache_async_set_cache_injects_default_in_memory_ttl():
|
|||
assert expiry <= after + 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dual_cache_redis_backfill_injects_default_in_memory_ttl():
|
||||
"""
|
||||
A Redis-hit backfill into the in-memory tier must honor
|
||||
default_in_memory_ttl the same way the write paths do. Without it, the
|
||||
backfilled entry falls to InMemoryCache's own default_ttl (600s), so a
|
||||
replica that primed a management object (e.g. a virtual key's auth blob)
|
||||
from Redis keeps serving it for 10 minutes after the object was updated
|
||||
and invalidated, instead of re-reading within the configured TTL.
|
||||
"""
|
||||
in_memory_cache = InMemoryCache(default_ttl=600)
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.async_get_cache = AsyncMock(return_value="redis_value")
|
||||
dual_cache = DualCache(
|
||||
in_memory_cache=in_memory_cache,
|
||||
redis_cache=redis_cache,
|
||||
default_in_memory_ttl=60,
|
||||
)
|
||||
|
||||
before = time.time()
|
||||
result = await dual_cache.async_get_cache(key="backfill_key")
|
||||
after = time.time()
|
||||
|
||||
assert result == "redis_value"
|
||||
expiry = in_memory_cache.ttl_dict["backfill_key"]
|
||||
assert expiry >= before + 60
|
||||
assert expiry <= after + 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dual_cache_batch_redis_backfill_injects_default_in_memory_ttl():
|
||||
"""async_batch_get_cache's Redis-to-memory backfill must honor
|
||||
default_in_memory_ttl, same as the single-key path."""
|
||||
in_memory_cache = InMemoryCache(default_ttl=600)
|
||||
mock_redis = MagicMock(spec=RedisCache)
|
||||
mock_redis.async_batch_get_cache = AsyncMock(
|
||||
return_value={"batch_backfill_key": "redis_value"}
|
||||
)
|
||||
dual_cache = DualCache(
|
||||
in_memory_cache=in_memory_cache,
|
||||
redis_cache=mock_redis,
|
||||
default_in_memory_ttl=60,
|
||||
)
|
||||
|
||||
before = time.time()
|
||||
result = await dual_cache.async_batch_get_cache(keys=["batch_backfill_key"])
|
||||
after = time.time()
|
||||
|
||||
assert result == ["redis_value"]
|
||||
expiry = in_memory_cache.ttl_dict["batch_backfill_key"]
|
||||
assert expiry >= before + 60
|
||||
assert expiry <= after + 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dual_cache_async_set_cache_respects_explicit_ttl():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -522,6 +522,49 @@ async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_e
|
|||
assert mock_prisma_client.get_data.await_count == 1
|
||||
|
||||
|
||||
def _fake_redis_cache():
|
||||
fake_redis = MagicMock()
|
||||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||||
fake_redis.async_set_cache = AsyncMock()
|
||||
fake_redis.async_set_cache_pipeline = AsyncMock()
|
||||
fake_redis.async_delete_cache = AsyncMock()
|
||||
return fake_redis
|
||||
|
||||
|
||||
class TestAuthCacheRedisWritePolicy:
|
||||
"""Redis auth-cache entries may only be written from fresh DB loads.
|
||||
|
||||
With ``enable_redis_auth_cache`` and multiple replicas, a pod that re-publishes
|
||||
a cache-derived key object to Redis can resurrect a stale auth blob after
|
||||
``/key/update`` or ``/key/delete`` already deleted it, so limit changes never
|
||||
propagate fleet-wide while traffic keeps refreshing the stale entry's TTL.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_db_load_publishes_to_redis(self):
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.get_data = AsyncMock(
|
||||
return_value=UserAPIKeyAuth(token="hashed-token-db")
|
||||
)
|
||||
|
||||
fake_redis = _fake_redis_cache()
|
||||
cache = UserApiKeyCache()
|
||||
cache.redis_cache = fake_redis
|
||||
|
||||
key_obj = await get_key_object(
|
||||
hashed_token="hashed-token-db",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=cache,
|
||||
)
|
||||
|
||||
assert key_obj.token == "hashed-token-db"
|
||||
fake_redis.async_set_cache.assert_awaited_once()
|
||||
assert (
|
||||
fake_redis.async_set_cache.await_args.kwargs.get("key")
|
||||
or fake_redis.async_set_cache.await_args.args[0]
|
||||
) == "hashed-token-db"
|
||||
|
||||
|
||||
def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values):
|
||||
"""Test generating CLI JWT token with default 24-hour expiration"""
|
||||
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
|
|
@ -3506,6 +3507,167 @@ async def test_centralized_common_checks_user_http_exception_isolates_to_user_on
|
|||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"key_org_id,team_org_id,expected_org_id",
|
||||
[
|
||||
(None, "org-from-team", "org-from-team"),
|
||||
("org-pinned-on-key", "org-from-team", "org-pinned-on-key"),
|
||||
(None, None, None),
|
||||
],
|
||||
)
|
||||
async def test_centralized_common_checks_backfills_org_id_from_team(key_org_id, team_org_id, expected_org_id):
|
||||
"""LIT-4688 regression: a key minted without an organization_id but attached
|
||||
to an org-linked team must leave auth with org_id set from the team, so the
|
||||
spend writer (which reads user_api_key_dict.org_id, no team fallback)
|
||||
credits the org and the org budget cap can actually trip. A key with an
|
||||
explicitly pinned org_id must win over the team's org."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
|
||||
token = UserAPIKeyAuth(api_key="sk-test", user_id="u", team_id="t1", org_id=key_org_id)
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
fetched_team = LiteLLM_TeamTableCachedObj(team_id="t1", organization_id=team_org_id)
|
||||
|
||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
org_id_seen_by_common_checks = []
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=fetched_team,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=lambda **kw: org_id_seen_by_common_checks.append(kw["valid_token"].org_id),
|
||||
) as mock_checks,
|
||||
):
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
mock_checks.assert_awaited_once()
|
||||
assert token.org_id == expected_org_id
|
||||
assert org_id_seen_by_common_checks == [expected_org_id]
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_session_token_org_backfilled_from_team(monkeypatch):
|
||||
"""LIT-4688 root cause: CLI session tokens (from /sso/cli/poll) are minted
|
||||
with a real team_id but no org_id, and their auth path decrypts the blob
|
||||
without the combined_view team join, so their spend never reached the org.
|
||||
The centralized-checks backfill must complete the credential from the team
|
||||
the same way the SQL view does for DB keys."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-lit4688")
|
||||
|
||||
cli_user = LiteLLM_UserTable(user_id="cli-user", user_role="internal_user", teams=["t-cli"], models=[])
|
||||
blob = ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id="t-cli", team_alias="cli-team")
|
||||
token = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(blob)
|
||||
assert token is not None
|
||||
assert token.is_session_token is True
|
||||
assert token.team_id == "t-cli"
|
||||
assert token.org_id is None
|
||||
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
org_linked_team = LiteLLM_TeamTableCachedObj(team_id="t-cli", organization_id="org-infoops")
|
||||
|
||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=org_linked_team,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
):
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
assert token.org_id == "org-infoops"
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_centralized_common_checks_org_backfill_survives_team_fetch_failure():
|
||||
"""When the team DB fetch fails, the token-derived fallback team carries no
|
||||
organization_id, so the backfill must leave org_id as None rather than
|
||||
crash or mis-attribute."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
token = UserAPIKeyAuth(api_key="sk-test", user_id="u", team_id="t1")
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
|
||||
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=Exception("DB down"),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.common_checks",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_checks,
|
||||
):
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/chat/completions",
|
||||
)
|
||||
|
||||
mock_checks.assert_awaited_once()
|
||||
assert token.org_id is None
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_master_key_auth_substitutes_alias_for_api_key():
|
||||
"""
|
||||
|
|
@ -4087,6 +4249,99 @@ async def test_auth_path_caches_team_object_under_canonical_team_id_key():
|
|||
assert cache.get_cache(key=None) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_does_not_rewrite_cached_key_object_back_into_cache():
|
||||
"""A cache-hit auth must not write the token back into the cache.
|
||||
|
||||
Re-writing on every auth let a replica holding a stale in-memory token
|
||||
republish it to shared Redis with a fresh TTL on each request, so
|
||||
/key/update and /key/delete never propagated across replicas or regional
|
||||
Redis while the key kept calling (stale auth re-cache feedback loop).
|
||||
Only the DB-load paths (IdentityStore._resolve_key / get_key_object) may
|
||||
populate the cache.
|
||||
"""
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.proxy_server import hash_token
|
||||
|
||||
api_key = "sk-lit-cached-key-no-rewrite"
|
||||
hashed_key = hash_token(api_key)
|
||||
|
||||
key_cache = UserApiKeyCache()
|
||||
stale_token = UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
token=hashed_key,
|
||||
metadata={"model_rpm_limit": {"gpt-5.4-mini": 3}},
|
||||
last_refreshed_at=1000.0,
|
||||
)
|
||||
await key_cache.async_set_cache(
|
||||
key=hashed_key, value=stale_token, model_type=UserAPIKeyAuth
|
||||
)
|
||||
|
||||
fetch_from_db = AsyncMock(
|
||||
side_effect=AssertionError("cache-hit auth must not touch the DB")
|
||||
)
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||
proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
|
||||
attrs = {
|
||||
"prisma_client": MagicMock(),
|
||||
"user_api_key_cache": key_cache,
|
||||
"proxy_logging_obj": proxy_logging_obj,
|
||||
"master_key": "sk-test-master",
|
||||
"general_settings": {},
|
||||
"llm_model_list": [],
|
||||
"llm_router": None,
|
||||
"open_telemetry_logger": None,
|
||||
"model_max_budget_limiter": MagicMock(),
|
||||
"user_custom_auth": None,
|
||||
"jwt_handler": None,
|
||||
"litellm_proxy_admin_name": "admin",
|
||||
}
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
request = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
with patch(
|
||||
"litellm.proxy.auth.resolvers.store._fetch_key_object_from_db_with_reconnect",
|
||||
fetch_from_db,
|
||||
):
|
||||
result = await _user_api_key_auth_builder(
|
||||
request=request,
|
||||
api_key=f"Bearer {api_key}",
|
||||
azure_api_key_header="",
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
request_data={},
|
||||
)
|
||||
pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
|
||||
if pending:
|
||||
await asyncio.wait(pending, timeout=5)
|
||||
|
||||
assert result.token == hashed_key
|
||||
fetch_from_db.assert_not_called()
|
||||
|
||||
cached_after = await key_cache.async_get_cache(
|
||||
key=hashed_key, model_type=UserAPIKeyAuth
|
||||
)
|
||||
assert cached_after is not None
|
||||
assert cached_after.last_refreshed_at == 1000.0
|
||||
assert cached_after.metadata == {"model_rpm_limit": {"gpt-5.4-mini": 3}}
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
class TestCheckKeyModelBudgetWithFallback:
|
||||
"""`_check_key_model_budget_with_fallback` must reroute a request to the
|
||||
first configured `budget_fallbacks` entry still within its own budget,
|
||||
|
|
@ -4347,3 +4602,84 @@ class TestCheckKeyModelBudgetWithFallback:
|
|||
|
||||
assert exc_info.value is original_error
|
||||
assert "model" not in request_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_temp_budget_increase_applied_for_cached_key():
|
||||
"""
|
||||
Regression for https://github.com/BerriAI/litellm/issues/25760
|
||||
|
||||
temp_budget_increase used to be applied only on the DB-fetch path, so a key
|
||||
served from cache kept its original max_budget and was wrongly blocked once
|
||||
spend crossed the original budget (but stayed under the effective budget).
|
||||
|
||||
Seed the auth cache with a key whose spend (5.0) exceeds its original
|
||||
max_budget (2.0) but is under the effective budget (2.0 + 100.0). The cache-hit
|
||||
request must not raise and the resolved token must carry max_budget == 102.0.
|
||||
|
||||
Resolving twice must yield 102.0 both times and leave the cached object at the
|
||||
original 2.0: the increase is derived per request, never compounded or persisted.
|
||||
"""
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from litellm.proxy.utils import hash_token
|
||||
|
||||
api_key = "sk-temp-budget-cache-regression"
|
||||
hashed_token = hash_token(api_key)
|
||||
expiry = (datetime.now() + timedelta(days=1)).isoformat()
|
||||
|
||||
cached_key = UserAPIKeyAuth(
|
||||
token=hashed_token,
|
||||
max_budget=2.0,
|
||||
spend=5.0,
|
||||
metadata={"temp_budget_increase": 100.0, "temp_budget_expiry": expiry},
|
||||
)
|
||||
|
||||
user_api_key_cache = DualCache()
|
||||
await _cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
user_api_key_obj=cached_key,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/v1/chat/completions"
|
||||
mock_request.method = "POST"
|
||||
mock_request.headers = {"authorization": f"Bearer {api_key}"}
|
||||
mock_request.query_params = {}
|
||||
mock_request.state = SimpleNamespace()
|
||||
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.budget_alerts = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-master"),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj),
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth._virtual_key_max_budget_alert_check",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
):
|
||||
results = tuple(
|
||||
[
|
||||
await _user_api_key_auth_builder(
|
||||
request=mock_request,
|
||||
api_key=f"Bearer {api_key}",
|
||||
azure_api_key_header="",
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
request_data={"model": "gpt-4o-mini"},
|
||||
)
|
||||
for _ in range(2)
|
||||
]
|
||||
)
|
||||
|
||||
assert all(result.max_budget == 102.0 for result in results)
|
||||
|
||||
cached_after = await user_api_key_cache.async_get_cache(key=hashed_token)
|
||||
assert cached_after.max_budget == 2.0
|
||||
|
|
|
|||
|
|
@ -2139,6 +2139,93 @@ class TestCLIKeyRegenerationFlow:
|
|||
assert not _is_valid_cli_sso_login_id("cli-test\x001234567890")
|
||||
assert not _is_valid_cli_sso_login_id("sk-test1234567890")
|
||||
|
||||
def test_cli_sso_flow_is_redis_authoritative_when_redis_attached(self):
|
||||
"""
|
||||
When Redis is attached, the CLI SSO flow must be read from and written to
|
||||
Redis directly, never the in-memory layer. Otherwise the worker that served
|
||||
/sso/cli/start keeps serving its stale in-memory flow and never sees the
|
||||
sso_complete/session_data update another worker wrote, which is exactly the
|
||||
multi-worker failure this fix targets.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
CLI_SSO_SESSION_TTL_SECONDS,
|
||||
_get_cli_sso_flow_cache_key,
|
||||
_get_cli_sso_flow_or_raise,
|
||||
_set_cli_sso_flow,
|
||||
)
|
||||
|
||||
login_id = "cli-redis_authoritative_1234567890"
|
||||
cache_key = _get_cli_sso_flow_cache_key(login_id)
|
||||
fresh_flow = {"poll_secret_hash": "fresh", "sso_complete": True}
|
||||
stale_flow = {"poll_secret_hash": "stale", "sso_complete": False}
|
||||
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.get_cache.return_value = fresh_flow
|
||||
cache = MagicMock()
|
||||
cache.redis_cache = redis_cache
|
||||
cache.get_cache.return_value = stale_flow
|
||||
|
||||
result = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cache)
|
||||
|
||||
assert result == fresh_flow
|
||||
redis_cache.get_cache.assert_called_once_with(key=cache_key)
|
||||
cache.get_cache.assert_not_called()
|
||||
|
||||
_set_cli_sso_flow(login_id=login_id, cache=cache, flow=fresh_flow)
|
||||
|
||||
redis_cache.set_cache.assert_called_once_with(
|
||||
key=cache_key, value=json.dumps(fresh_flow), ttl=CLI_SSO_SESSION_TTL_SECONDS
|
||||
)
|
||||
cache.set_cache.assert_not_called()
|
||||
|
||||
def test_cli_sso_flow_with_enum_survives_redis_round_trip(self):
|
||||
"""
|
||||
RedisCache stores values via str(value) and reads them back through
|
||||
json.loads/ast.literal_eval. A raw flow dict containing a Python enum
|
||||
(session_data.user_role after the SSO callback) produces an unparseable
|
||||
repr, so every worker reading the completed flow from Redis got a
|
||||
SyntaxError and returned 400 "session not found". The flow must survive
|
||||
a real Redis serialization round trip.
|
||||
"""
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_get_cli_sso_flow_or_raise,
|
||||
_set_cli_sso_flow,
|
||||
)
|
||||
|
||||
login_id = "cli-enum_round_trip_1234567890"
|
||||
completed_flow = {
|
||||
"poll_secret_hash": "hash",
|
||||
"sso_complete": True,
|
||||
"user_code_verified": False,
|
||||
"session_data": {
|
||||
"user_id": "user-1",
|
||||
"user_role": LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||
"models": [],
|
||||
"teams": ["team-1"],
|
||||
"team_details": [{"team_id": "team-1", "team_alias": "alias"}],
|
||||
},
|
||||
}
|
||||
|
||||
redis_store: dict = {}
|
||||
redis_cache = MagicMock()
|
||||
redis_cache.set_cache.side_effect = lambda key, value, ttl: redis_store.__setitem__(
|
||||
key, str(value).encode("utf-8")
|
||||
)
|
||||
redis_cache.get_cache.side_effect = lambda key: RedisCache._get_cache_logic(
|
||||
MagicMock(), redis_store.get(key)
|
||||
)
|
||||
cache = MagicMock()
|
||||
cache.redis_cache = redis_cache
|
||||
|
||||
_set_cli_sso_flow(login_id=login_id, cache=cache, flow=completed_flow)
|
||||
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cache)
|
||||
|
||||
assert flow["sso_complete"] is True
|
||||
assert flow["session_data"]["user_role"] == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value
|
||||
assert flow["session_data"]["team_details"] == [{"team_id": "team-1", "team_alias": "alias"}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_sso_start_creates_bound_flow(self):
|
||||
"""Test CLI SSO start creates a polling secret bound flow"""
|
||||
|
|
@ -2151,10 +2238,13 @@ class TestCLIKeyRegenerationFlow:
|
|||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.client = SimpleNamespace(host="127.0.0.1")
|
||||
mock_request.headers = {}
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.increment_cache.return_value = 1
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
result = await cli_sso_start(request=mock_request)
|
||||
|
||||
assert result["login_id"].startswith("cli-")
|
||||
|
|
@ -2182,10 +2272,13 @@ class TestCLIKeyRegenerationFlow:
|
|||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.client = SimpleNamespace(host="127.0.0.1")
|
||||
mock_request.headers = {}
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.increment_cache.return_value = 31
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await cli_sso_start(request=mock_request)
|
||||
|
||||
|
|
@ -2204,7 +2297,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
mock_request.client = SimpleNamespace(host="127.0.0.1")
|
||||
mock_request.headers = {}
|
||||
mock_request.base_url = "https://proxy.example.com/"
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.increment_cache.return_value = 1
|
||||
|
||||
with (
|
||||
|
|
@ -2238,7 +2331,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
mock_request.client = SimpleNamespace(host="127.0.0.1")
|
||||
mock_request.headers = {}
|
||||
mock_request.base_url = "https://proxy.example.com/"
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.increment_cache.return_value = 1
|
||||
|
||||
with (
|
||||
|
|
@ -2272,7 +2365,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.example.com/"
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {"poll_secret_hash": "h"}
|
||||
|
||||
async def drive(enabled: bool):
|
||||
|
|
@ -2281,6 +2374,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler",
|
||||
None,
|
||||
|
|
@ -2448,7 +2542,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
)
|
||||
mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"}
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": "poll-secret-hash",
|
||||
"user_code_hash": "user-code-hash",
|
||||
|
|
@ -2467,6 +2561,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
result = await cli_sso_callback(
|
||||
request=mock_request,
|
||||
|
|
@ -2491,7 +2586,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
mock_request.body = AsyncMock(
|
||||
return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token"
|
||||
)
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"user_code_hash": _hash_cli_sso_secret(
|
||||
|
|
@ -2505,6 +2600,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch(
|
||||
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
|
||||
return_value="<html>Success</html>",
|
||||
|
|
@ -2529,7 +2625,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.body = AsyncMock(return_value=b"user_code=ABCD-EFGH")
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"user_code_hash": _hash_cli_sso_secret(
|
||||
|
|
@ -2541,7 +2637,10 @@ class TestCLIKeyRegenerationFlow:
|
|||
"session_data": {"user_id": "test-user-123"},
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await cli_sso_complete(
|
||||
request=mock_request, login_id="cli-session-4567890"
|
||||
|
|
@ -2563,7 +2662,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
mock_request.body = AsyncMock(
|
||||
return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token"
|
||||
)
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"user_code_hash": _hash_cli_sso_secret(
|
||||
|
|
@ -2574,7 +2673,10 @@ class TestCLIKeyRegenerationFlow:
|
|||
"session_data": None,
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await cli_sso_complete(
|
||||
request=mock_request, login_id="cli-session-4567890"
|
||||
|
|
@ -2610,7 +2712,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"}
|
||||
|
||||
# Mock cache
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": "poll-secret-hash",
|
||||
"user_code_hash": "user-code-hash",
|
||||
|
|
@ -2632,6 +2734,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch(
|
||||
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
|
||||
return_value="<html>Success</html>",
|
||||
|
|
@ -2692,7 +2795,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
}
|
||||
|
||||
# Mock cache
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -2700,7 +2803,10 @@ class TestCLIKeyRegenerationFlow:
|
|||
"session_data": session_data,
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
# Act - First poll without team_id
|
||||
result = await cli_poll_key(
|
||||
key_id=session_key,
|
||||
|
|
@ -2726,7 +2832,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
cli_poll_key,
|
||||
)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -2739,7 +2845,10 @@ class TestCLIKeyRegenerationFlow:
|
|||
},
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await cli_poll_key(key_id="cli-session-789123", team_id=None)
|
||||
|
||||
|
|
@ -2753,7 +2862,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
cli_poll_key,
|
||||
)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -2766,7 +2875,10 @@ class TestCLIKeyRegenerationFlow:
|
|||
},
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
result = await cli_poll_key(
|
||||
key_id="cli-session-789123",
|
||||
team_id=None,
|
||||
|
|
@ -2932,7 +3044,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
)
|
||||
|
||||
# Mock cache
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -2944,6 +3056,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.prisma_client"),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
|
||||
|
|
@ -3007,7 +3120,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
models=["gpt-4"],
|
||||
max_budget=100.0,
|
||||
)
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -3018,6 +3131,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.prisma_client"),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
|
||||
|
|
@ -3069,7 +3183,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
max_budget=None,
|
||||
)
|
||||
mock_team = LiteLLM_TeamTableCachedObj(team_id="team-x", max_budget=None)
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -3080,6 +3194,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.prisma_client"),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
|
||||
|
|
@ -4020,7 +4135,7 @@ class TestPKCEFunctionality:
|
|||
mock_request.query_params = {"state": test_state}
|
||||
|
||||
# Mock cache with async methods — use dict format (primary path)
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
test_code_verifier = "test_code_verifier_abc123xyz"
|
||||
mock_cache.async_get_cache = AsyncMock(
|
||||
return_value={"code_verifier": test_code_verifier}
|
||||
|
|
@ -4071,7 +4186,7 @@ class TestPKCEFunctionality:
|
|||
mock_sso.__exit__ = MagicMock(return_value=False)
|
||||
|
||||
test_state = "test456"
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
|
|
@ -4595,7 +4710,7 @@ class TestPKCEFunctionality:
|
|||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
|
@ -4721,7 +4836,7 @@ class TestPKCEFunctionality:
|
|||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
# Cache returns an integer — unexpected format
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=12345)
|
||||
mock_cache.async_delete_cache = AsyncMock()
|
||||
|
||||
|
|
@ -4763,7 +4878,7 @@ class TestPKCEFunctionality:
|
|||
|
||||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
|
@ -4851,7 +4966,7 @@ class TestPKCEFunctionality:
|
|||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
# Cache returns an integer — unexpected format
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=12345)
|
||||
mock_cache.async_delete_cache = AsyncMock()
|
||||
|
||||
|
|
@ -4903,7 +5018,7 @@ class TestPKCEFunctionality:
|
|||
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
|
||||
|
||||
legacy_verifier = "legacy_plain_string_verifier_abc123"
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=legacy_verifier)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
|
|
@ -6187,7 +6302,7 @@ class TestCliSsoAttributionMetadata:
|
|||
provider="generic",
|
||||
team_ids=[],
|
||||
)
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": "poll-secret-hash",
|
||||
"user_code_hash": "user-code-hash",
|
||||
|
|
@ -6204,6 +6319,7 @@ class TestCliSsoAttributionMetadata:
|
|||
),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.user_custom_sso", None),
|
||||
):
|
||||
await ui_sso.cli_sso_callback(
|
||||
|
|
@ -6228,7 +6344,7 @@ class TestCliSsoAttributionMetadata:
|
|||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "http://internal-proxy.local/"
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": "poll-secret-hash",
|
||||
"user_code_hash": "user-code-hash",
|
||||
|
|
@ -6251,6 +6367,7 @@ class TestCliSsoAttributionMetadata:
|
|||
) as get_user_info_mock,
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.user_custom_sso", None),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
|
|
@ -6297,7 +6414,7 @@ class TestCliSsoAttributionMetadata:
|
|||
"user_id": "test-user-123",
|
||||
"employment_type": "contractor",
|
||||
}
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": "poll-secret-hash",
|
||||
"user_code_hash": "user-code-hash",
|
||||
|
|
@ -6325,6 +6442,7 @@ class TestCliSsoAttributionMetadata:
|
|||
),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.user_custom_sso", None),
|
||||
patch(
|
||||
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
|
||||
|
|
@ -6366,7 +6484,7 @@ class TestCliSsoAttributionMetadata:
|
|||
"org": {"cost_center": "CC-42"},
|
||||
},
|
||||
}
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -6374,7 +6492,10 @@ class TestCliSsoAttributionMetadata:
|
|||
"session_data": session_data,
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
):
|
||||
result = await cli_poll_key(
|
||||
key_id=session_key,
|
||||
team_id=None,
|
||||
|
|
@ -7225,7 +7346,7 @@ async def test_cli_poll_key_tolerates_missing_user_row():
|
|||
"models": ["gpt-4"],
|
||||
}
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache = MagicMock(redis_cache=None)
|
||||
mock_cache.get_cache.return_value = {
|
||||
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
|
||||
"sso_complete": True,
|
||||
|
|
@ -7237,6 +7358,7 @@ async def test_cli_poll_key_tolerates_missing_user_row():
|
|||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache),
|
||||
patch("litellm.proxy.proxy_server.prisma_client"),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token",
|
||||
|
|
|
|||
|
|
@ -4490,6 +4490,128 @@ async def test_update_cache_pipeline_honors_user_api_key_cache_ttl():
|
|||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_tracking_never_writes_the_auth_object_back():
|
||||
"""Spend tracking must never write the auth object back into the cache.
|
||||
|
||||
Writing the mutated auth object back after every priced request let a
|
||||
stale copy be re-published with a fresh TTL: to shared Redis it defeated
|
||||
/key/update and /key/delete across replicas, and even a local-only write
|
||||
could race an invalidation and resurrect a revoked key on this worker.
|
||||
Spend is tracked through the spend:key:* counters, so the auth object is
|
||||
only ever written by the DB-load paths.
|
||||
"""
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
original_cache = litellm.proxy.proxy_server.user_api_key_cache
|
||||
cache = UserApiKeyCache()
|
||||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
|
||||
try:
|
||||
hashed_token = "spend-tracking-no-writeback-token"
|
||||
await cache.async_set_cache(
|
||||
key=hashed_token,
|
||||
value=UserAPIKeyAuth(token=hashed_token, spend=1.0),
|
||||
model_type=UserAPIKeyAuth,
|
||||
)
|
||||
with (
|
||||
patch.object(
|
||||
cache, "async_set_cache_pipeline", new=AsyncMock()
|
||||
) as mock_pipeline,
|
||||
patch.object(cache, "async_set_cache", new=AsyncMock()) as mock_set,
|
||||
):
|
||||
await litellm.proxy.proxy_server.update_cache(
|
||||
token=hashed_token,
|
||||
user_id=None,
|
||||
end_user_id=None,
|
||||
team_id=None,
|
||||
response_cost=5.0,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
pending = [
|
||||
t for t in asyncio.all_tasks() if t is not asyncio.current_task()
|
||||
]
|
||||
if pending:
|
||||
await asyncio.wait(pending, timeout=5)
|
||||
|
||||
key_pipeline_writes = [
|
||||
call
|
||||
for call in mock_pipeline.call_args_list
|
||||
if any(k == hashed_token for k, _ in call.kwargs["cache_list"])
|
||||
]
|
||||
assert key_pipeline_writes == []
|
||||
mock_set.assert_not_called()
|
||||
finally:
|
||||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_cache_global_proxy_spend_scalar_stays_shared():
|
||||
"""
|
||||
The proxy-wide spend estimate must keep flowing to Redis when the spend
|
||||
writeback goes per-pod: the global max_budget check reads the
|
||||
``{litellm_proxy_admin_name}:spend`` cache entry between authoritative DB
|
||||
reloads, so keeping it pod-local would let traffic spread across replicas
|
||||
exceed the proxy budget by roughly a factor of the replica count within a
|
||||
cache TTL. Sharing this scalar is safe because it carries no limits or
|
||||
permissions, so it cannot resurrect an invalidated auth blob.
|
||||
"""
|
||||
from litellm.caching.caching import DualCache
|
||||
|
||||
admin_name = litellm.proxy.proxy_server.litellm_proxy_admin_name
|
||||
global_key = "{}:spend".format(admin_name)
|
||||
|
||||
async def fake_get(key, **kwargs):
|
||||
if key == "user-lit":
|
||||
return {"user_id": "user-lit", "spend": 1.0}
|
||||
if key == global_key:
|
||||
return 10.0
|
||||
return None
|
||||
|
||||
original_cache = litellm.proxy.proxy_server.user_api_key_cache
|
||||
cache = DualCache(default_in_memory_ttl=300)
|
||||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
|
||||
try:
|
||||
with patch.object(
|
||||
cache, "async_get_cache", new=AsyncMock(side_effect=fake_get)
|
||||
):
|
||||
with patch.object(
|
||||
cache, "async_set_cache_pipeline", new=AsyncMock()
|
||||
) as mock_set_cache:
|
||||
await litellm.proxy.proxy_server.update_cache(
|
||||
token=None,
|
||||
user_id="user-lit",
|
||||
end_user_id=None,
|
||||
team_id=None,
|
||||
response_cost=5.0,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
|
||||
pending = [
|
||||
t for t in asyncio.all_tasks() if t is not asyncio.current_task()
|
||||
]
|
||||
if pending:
|
||||
await asyncio.wait(pending, timeout=5)
|
||||
|
||||
calls = mock_set_cache.await_args_list
|
||||
local_keys = [
|
||||
k
|
||||
for c in calls
|
||||
if c.kwargs.get("local_only") is True
|
||||
for k, _ in c.kwargs["cache_list"]
|
||||
]
|
||||
shared_keys = [
|
||||
k
|
||||
for c in calls
|
||||
if c.kwargs.get("local_only") is not True
|
||||
for k, _ in c.kwargs["cache_list"]
|
||||
]
|
||||
assert "user-lit" in local_keys
|
||||
assert global_key not in local_keys
|
||||
assert shared_keys == [global_key]
|
||||
finally:
|
||||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_sso_settings_in_db():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -54,8 +54,14 @@ def _patched_init_cache(litellm_settings: dict, cache_params: dict):
|
|||
_FakeRedisCache (passes the isinstance guard in _init_cache).
|
||||
3. Extracts enable_redis_auth_cache from litellm_settings and passes it
|
||||
as the second argument to _init_cache (matching production behaviour).
|
||||
4. Yields (user_api_key_cache, spend_counter_cache) after calling
|
||||
_init_cache, then restores everything.
|
||||
4. Yields (user_api_key_cache, spend_counter_cache, cli_sso_session_cache)
|
||||
after calling _init_cache, then restores everything.
|
||||
|
||||
_init_cache also writes three globals this helper does not patch:
|
||||
``litellm.cache``, ``ps.redis_usage_cache`` and
|
||||
``litellm_config_cache.redis_cache``. They are saved and restored here so a
|
||||
_FakeRedisCache never outlives this context and poisons later test files in
|
||||
the same pytest session.
|
||||
"""
|
||||
fake_redis = _FakeRedisCache()
|
||||
|
||||
|
|
@ -64,19 +70,30 @@ def _patched_init_cache(litellm_settings: dict, cache_params: dict):
|
|||
|
||||
fresh_user_cache = DualCache()
|
||||
fresh_spend_cache = DualCache()
|
||||
fresh_cli_sso_cache = DualCache()
|
||||
|
||||
enable_redis_auth_cache = litellm_settings.get("enable_redis_auth_cache", False)
|
||||
|
||||
prev_litellm_cache = litellm.cache
|
||||
prev_redis_usage_cache = ps.redis_usage_cache
|
||||
prev_config_cache_redis = ps.litellm_config_cache.redis_cache
|
||||
|
||||
with (
|
||||
patch.object(ps, "user_api_key_cache", fresh_user_cache),
|
||||
patch.object(ps, "spend_counter_cache", fresh_spend_cache),
|
||||
patch.object(ps, "cli_sso_session_cache", fresh_cli_sso_cache),
|
||||
patch.object(ps, "llm_router", None),
|
||||
# Cache is locally imported inside _init_cache: patch it at source.
|
||||
patch("litellm.Cache", return_value=mock_litellm_cache),
|
||||
):
|
||||
litellm.cache = None
|
||||
ps.ProxyConfig()._init_cache(cache_params, enable_redis_auth_cache)
|
||||
yield fresh_user_cache, fresh_spend_cache
|
||||
try:
|
||||
ps.ProxyConfig()._init_cache(cache_params, enable_redis_auth_cache)
|
||||
yield fresh_user_cache, fresh_spend_cache, fresh_cli_sso_cache
|
||||
finally:
|
||||
litellm.cache = prev_litellm_cache
|
||||
ps.redis_usage_cache = prev_redis_usage_cache
|
||||
ps.litellm_config_cache.redis_cache = prev_config_cache_redis
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -90,7 +107,7 @@ class TestRedisAuthCacheFlag:
|
|||
with _patched_init_cache(
|
||||
litellm_settings={"enable_redis_auth_cache": True},
|
||||
cache_params={"type": "redis", "host": "localhost", "port": 6379},
|
||||
) as (user_cache, _):
|
||||
) as (user_cache, _, _cli_sso_cache):
|
||||
assert user_cache.redis_cache is not None, (
|
||||
"Redis should be attached to user_api_key_cache when "
|
||||
"enable_redis_auth_cache=True"
|
||||
|
|
@ -101,7 +118,7 @@ class TestRedisAuthCacheFlag:
|
|||
with _patched_init_cache(
|
||||
litellm_settings={"enable_redis_auth_cache": False},
|
||||
cache_params={"type": "redis", "host": "localhost", "port": 6379},
|
||||
) as (user_cache, _):
|
||||
) as (user_cache, _, _cli_sso_cache):
|
||||
assert user_cache.redis_cache is None, (
|
||||
"user_api_key_cache must remain in-memory-only when "
|
||||
"enable_redis_auth_cache=False"
|
||||
|
|
@ -112,7 +129,7 @@ class TestRedisAuthCacheFlag:
|
|||
with _patched_init_cache(
|
||||
litellm_settings={},
|
||||
cache_params={"type": "redis", "host": "localhost", "port": 6379},
|
||||
) as (user_cache, _):
|
||||
) as (user_cache, _, _cli_sso_cache):
|
||||
assert user_cache.redis_cache is None, (
|
||||
"user_api_key_cache must remain in-memory-only when "
|
||||
"enable_redis_auth_cache is absent from litellm_settings"
|
||||
|
|
@ -129,7 +146,7 @@ class TestRedisAuthCacheFlag:
|
|||
with _patched_init_cache(
|
||||
litellm_settings=ls,
|
||||
cache_params={"type": "redis", "host": "localhost", "port": 6379},
|
||||
) as (_, spend_cache):
|
||||
) as (_, spend_cache, _cli_sso_cache):
|
||||
assert spend_cache.redis_cache is not None, (
|
||||
f"spend_counter_cache must always get Redis "
|
||||
f"(enable_redis_auth_cache={flag_value!r})"
|
||||
|
|
@ -140,6 +157,28 @@ class TestRedisAuthCacheFlag:
|
|||
with _patched_init_cache(
|
||||
litellm_settings={"enable_redis_auth_cache": False},
|
||||
cache_params={"type": "redis", "host": "localhost", "port": 6379},
|
||||
) as (user_cache, spend_cache):
|
||||
) as (user_cache, spend_cache, _cli_sso_cache):
|
||||
assert spend_cache.redis_cache is not None
|
||||
assert user_cache.redis_cache is None
|
||||
|
||||
def test_cli_sso_session_cache_always_gets_redis_regardless_of_flag(self):
|
||||
"""
|
||||
cli_sso_session_cache must receive Redis regardless of the auth-cache
|
||||
flag so that `lite login` works on multi-worker deployments without
|
||||
enable_redis_auth_cache (regression for the CLI SSO "Invalid CLI login
|
||||
session" bug)
|
||||
"""
|
||||
for flag_value in (True, False, None):
|
||||
ls = (
|
||||
{"enable_redis_auth_cache": flag_value}
|
||||
if flag_value is not None
|
||||
else {}
|
||||
)
|
||||
with _patched_init_cache(
|
||||
litellm_settings=ls,
|
||||
cache_params={"type": "redis", "host": "localhost", "port": 6379},
|
||||
) as (_, _, cli_sso_cache):
|
||||
assert cli_sso_cache.redis_cache is not None, (
|
||||
f"cli_sso_session_cache must always get Redis "
|
||||
f"(enable_redis_auth_cache={flag_value!r})"
|
||||
)
|
||||
|
|
|
|||
22
uv.lock
generated
22
uv.lock
generated
|
|
@ -10,7 +10,7 @@ resolution-markers = [
|
|||
]
|
||||
|
||||
[options]
|
||||
exclude-newer = "2026-07-15T20:33:58.442813Z"
|
||||
exclude-newer = "2026-07-23T00:33:07.241959Z"
|
||||
exclude-newer-span = "P3D"
|
||||
|
||||
[manifest]
|
||||
|
|
@ -2093,14 +2093,14 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "gitpython"
|
||||
version = "3.1.50"
|
||||
version = "3.1.54"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "gitdb" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/33/f6/354ae6491228b5eb40e10d89c4d13c651fe1cf7556e35ebdded50cff57ce/gitpython-3.1.50.tar.gz", hash = "sha256:80da2d12504d52e1f998772dc5baf6e553f8d2fcfe1fcc226c9d9a2ee3372dcc", size = 219798, upload-time = "2026-05-06T04:01:26.571Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/5e/d5/3da0b92033887033f4c27f2dd109a303c4ca62813c7b3bb2511edb4777de/gitpython-3.1.54.tar.gz", hash = "sha256:53f2085e24a2cda300eed7c3fc5f1559ae289634b725e98acaf4791940247aa0", size = 225076, upload-time = "2026-07-22T04:08:51.403Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/20/7a/1c6e3562dfd8950adbb11ffbc65d21e7c89d01a6e4f137fa981056de25c5/gitpython-3.1.50-py3-none-any.whl", hash = "sha256:d352abe2908d07355014abdd21ddf798c2a961469239afec4962e9da884858f9", size = 212507, upload-time = "2026-05-06T04:01:23.799Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/d1/b9/876f442a28df5c068ca69b0122d5c35e65fd2d2fa9992ea5cb5944ea00a6/gitpython-3.1.54-py3-none-any.whl", hash = "sha256:b90d7b3d9bc0238681d24369130826f0dcdb0ceaa45db67cf1d4ffa4c302dedf", size = 216575, upload-time = "2026-07-22T04:08:50.05Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -3732,7 +3732,7 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "litellm"
|
||||
version = "1.93.0"
|
||||
version = "1.93.1"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "aiohttp" },
|
||||
|
|
@ -6454,11 +6454,11 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "pyasn1"
|
||||
version = "0.6.3"
|
||||
version = "0.6.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/5c/5f/6583902b6f79b399c9c40674ac384fd9cd77805f9e6205075f828ef11fb2/pyasn1-0.6.3.tar.gz", hash = "sha256:697a8ecd6d98891189184ca1fa05d1bb00e2f84b5977c481452050549c8a72cf", size = 148685, upload-time = "2026-03-17T01:06:53.382Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/a4/9a/23310166d960def5897e91fe20e5b724601b02a22e84ba1f94232c0b7f67/pyasn1-0.6.4.tar.gz", hash = "sha256:9c447d8431c947fe4c8febc4ed9e760bc29011a5b01e5c74b67025bd9fb8ce81", size = 151262, upload-time = "2026-07-09T01:12:33.988Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/5d/a0/7d793dce3fa811fe047d6ae2431c672364b462850c6235ae306c0efd025f/pyasn1-0.6.3-py3-none-any.whl", hash = "sha256:a80184d120f0864a52a073acc6fc642847d0be408e7c7252f31390c0f4eadcde", size = 83997, upload-time = "2026-03-17T01:06:52.036Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/9a/3b/6163796d69c3977d1e4287bea4a6979161cbbdd170ebb430511e8e1999ce/pyasn1-0.6.4-py3-none-any.whl", hash = "sha256:deda9277cfd454080ec40b207fb6df82206a3a2688735233cdcd8d3d565f088b", size = 84410, upload-time = "2026-07-09T01:12:32.92Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -6778,14 +6778,14 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "pypdf"
|
||||
version = "6.13.3"
|
||||
version = "6.14.2"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "typing-extensions", marker = "python_full_version < '3.11'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/17/18/9947cc201af9ccf76720fd3347bf4f70eb882ce3fcf4cb05f7443e4cf871/pypdf-6.13.3.tar.gz", hash = "sha256:f3cb822769725f1bac658c406cfc9460399043f3750c2d3e4650e0a85eacabd7", size = 6484063, upload-time = "2026-06-17T15:22:00.898Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/03/72/7dfd5ff1c9c37de97a731701f51af091325f123d9d4270361c9c69e4431f/pypdf-6.14.2.tar.gz", hash = "sha256:7873f502fe4385e79539b21d872392dc0c4e3714327c15881cbc7fbfd1f95b25", size = 6491182, upload-time = "2026-06-23T14:18:30.859Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/94/56/2967e621598987905fb8cdfadd8f8de6b5c68c9351f0523c4df8409f28f1/pypdf-6.13.3-py3-none-any.whl", hash = "sha256:c6e3f86afb625791510b02ad5480e94b63970bb957df75d44657c282ecc52224", size = 347288, upload-time = "2026-06-17T15:21:59.512Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/49/e6/136aa8993a2ae7214e0b0ef2edaa0d2e08d1d4e4982635b08a835ff31ec8/pypdf-6.14.2-py3-none-any.whl", hash = "sha256:3f07891af76dc002657e04993ab9b4de81de29f9013b9761d0b7968bff12e946", size = 349514, upload-time = "2026-06-23T14:18:28.867Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue