chore: merge origin/main into litellm_langfuse_otel_end_user_id

This commit is contained in:
yucheng 2026-09-23 13:51:16 +00:00
commit e4b5f26a77
35 changed files with 2169 additions and 234 deletions

View file

@ -1,186 +0,0 @@
name: Create Release
on:
workflow_dispatch:
inputs:
tag:
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0-dev.2, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
required: true
type: string
commit_hash:
description: "Full 40-char commit SHA to target"
required: true
type: string
permissions: {}
jobs:
release:
name: Create Release
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Validate inputs
env:
TAG: ${{ inputs.tag }}
COMMIT_HASH: ${{ inputs.commit_hash }}
run: |
if ! echo "${COMMIT_HASH}" | grep -qE '^[0-9a-f]{40}$'; then
echo "::error::commit_hash must be a full 40-character commit SHA"
exit 1
fi
if ! echo "${TAG}" | grep -qE '^v?[0-9]+\.[0-9]+\.[0-9]+'; then
echo "::error::tag must start with X.Y.Z (optional leading v), e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, or v1.83.10-stable"
exit 1
fi
- name: Create release
env:
TAG: ${{ inputs.tag }}
COMMIT_HASH: ${{ inputs.commit_hash }}
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
with:
script: |
const tag = process.env.TAG;
const commitHash = process.env.COMMIT_HASH;
// Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases.
// Accept both PEP 440 (`.dev`) and SemVer (`-dev`) separators so tags
// like `1.84.0.dev2` and `1.84.0-dev.2` are both detected.
// PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]`
// are stable maintenance releases, not pre-releases.
const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag);
// A stable release should only claim the repo "latest" badge when its
// version is >= the current latest. Otherwise a backport (e.g. 1.84.6)
// would steal "latest" from a newer line (e.g. 1.88.1).
const versionKey = (rawTag) => {
const m = String(rawTag).match(/^v?(\d+)\.(\d+)\.(\d+)/);
if (!m) return null;
const maintenance = String(rawTag).match(/(?:\.post|\.patch\.)(\d+)/i);
return [Number(m[1]), Number(m[2]), Number(m[3]), maintenance ? Number(maintenance[1]) : 0];
};
const isAtLeast = (a, b) => {
for (let i = 0; i < a.length; i++) {
if (a[i] !== b[i]) return a[i] > b[i];
}
return true;
};
const cosignSection = [
`## Verify Docker Image Signature`,
``,
`All LiteLLM Docker images are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit \`0112e53\`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).`,
``,
`**Verify using the pinned commit hash (recommended):**`,
``,
`A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:`,
``,
'```bash',
`cosign verify \\`,
` --key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \\`,
` ghcr.io/berriai/litellm:${tag}`,
'```',
``,
`**Verify using the release tag (convenience):**`,
``,
`Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:`,
``,
'```bash',
`cosign verify \\`,
` --key https://raw.githubusercontent.com/BerriAI/litellm/${tag}/cosign.pub \\`,
` ghcr.io/berriai/litellm:${tag}`,
'```',
``,
`Expected output:`,
``,
'```',
`The following checks were performed on each of these signatures:`,
` - The cosign claims were validated`,
` - The signatures were verified against the specified public key`,
'```',
``,
`---`,
``,
].join('\n');
try {
let makeLatest = "false";
const newVersion = versionKey(tag);
if (!isPrerelease && newVersion) {
let latestVersion = null;
try {
const latest = await github.rest.repos.getLatestRelease({
owner: context.repo.owner,
repo: context.repo.repo,
});
latestVersion = versionKey(latest.data.tag_name);
} catch (error) {
if (error.status !== 404) throw error;
}
makeLatest = (!latestVersion || isAtLeast(newVersion, latestVersion)) ? "true" : "false";
}
try {
await github.rest.git.createRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `refs/tags/${tag}`,
sha: commitHash,
});
} catch (error) {
if (error.status !== 422) throw error;
const existing = await github.rest.git.getRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `tags/${tag}`,
});
if (existing.data.object.sha !== commitHash) {
throw new Error(`Tag ${tag} already exists at ${existing.data.object.sha}, expected ${commitHash}`);
}
}
const response = await github.rest.repos.createRelease({
draft: true,
generate_release_notes: true,
name: tag,
owner: context.repo.owner,
prerelease: isPrerelease,
repo: context.repo.repo,
tag_name: tag,
});
const updatedBody = cosignSection + (response.data.body ?? '');
await github.rest.repos.updateRelease({
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
body: updatedBody,
draft: false,
});
if (!isPrerelease) {
await github.rest.repos.updateRelease({
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
make_latest: makeLatest,
});
}
} catch (error) {
core.setFailed(error.message);
}
create-branch:
name: Create Release Branch
needs: release
permissions:
contents: write
uses: ./.github/workflows/create-release-branch.yml
with:
tag: ${{ inputs.tag }}
commit_hash: ${{ inputs.commit_hash }}

View file

@ -26,6 +26,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/v2/login",
"/v3/login",
"/logout",
"/session/logout",
"/token",
"/onboarding/",
"/audit",

View file

@ -9,7 +9,7 @@ import sys
from collections.abc import Iterator
from datetime import datetime
from logging import Formatter
from typing import Any, Final, TextIO
from typing import Final, TextIO
from urllib.parse import unquote
import litellm
@ -672,13 +672,13 @@ def _try_parse_json_message(message: str) -> dict[str, object] | None:
msg_stripped: Final = message.strip()
if not (msg_stripped.startswith("{") or msg_stripped.startswith("[")):
return None
parsed: Final = safe_json_loads(message, default=None)
parsed: Final[object] = safe_json_loads(message, default=None)
if parsed is None or not isinstance(parsed, dict):
return None
return parsed
def _try_parse_embedded_python_dict(message: str) -> dict[str, Any] | None:
def _try_parse_embedded_python_dict(message: str) -> dict[str, object] | None:
"""
Try to find and parse a Python dict repr (e.g. str(d) or repr(d)) embedded in
the message. Handles patterns like:
@ -702,7 +702,7 @@ def _try_parse_embedded_python_dict(message: str) -> dict[str, Any] | None:
if depth == 0:
substr = message[start : j + 1]
try:
result = ast.literal_eval(substr)
result: object = ast.literal_eval(substr)
if isinstance(result, dict) and len(result) > 0:
return result
except (ValueError, SyntaxError, TypeError):

View file

@ -561,6 +561,7 @@ class OpenTelemetryV2(CustomLogger):
time_to_first_chunk_seconds=call.time_to_first_chunk_seconds,
request_route=request_root_http_route(),
trace=call.trace,
session_id=call.session_id,
)
end_time_ns: Final = to_ns(end_time)
if carrier is not None and carrier.span is not None:

View file

@ -43,6 +43,7 @@ class GenAIMapper:
GenAI.OPERATION_NAME: lambda d: d.operation.value,
GenAI.PROVIDER_NAME: lambda d: d.provider or None,
GenAI.OUTPUT_TYPE: lambda d: d.output_type.value if d.output_type else None,
GenAI.CONVERSATION_ID: lambda d: d.session_id,
GenAI.REQUEST_MODEL: lambda d: d.request_model or None,
GenAI.REQUEST_TEMPERATURE: lambda d: d.request_params.temperature,
GenAI.REQUEST_TOP_P: lambda d: d.request_params.top_p,

View file

@ -41,7 +41,7 @@ from dataclasses import dataclass, field
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, cast
from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL
from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL, SESSION_ID_GENERATED_METADATA_KEY
from litellm.integrations.otel.model.semconv import resolve_operation
from litellm.integrations.otel.model.trace_controls import TraceControls, caller_trace_controls
from litellm.integrations.otel.model.utils import as_str, as_str_mapping, to_seconds
@ -226,6 +226,7 @@ class LLMCallEvent:
provisional_span_name: str
time_to_first_chunk_seconds: float | None
trace: TraceControls
session_id: str | None
@classmethod
def from_dict(cls, kwargs: Mapping[str, object]) -> LLMCallEvent:
@ -233,6 +234,7 @@ class LLMCallEvent:
payload: Final = cast("StandardLoggingPayload", raw_payload) if raw_payload else None
operation: Final = resolve_operation(as_str(kwargs.get("call_type")))
model: Final = as_str(kwargs.get("model")) or ""
trace: Final = caller_trace_controls(kwargs)
return cls(
call_id=_call_id(payload, kwargs),
payload=payload,
@ -242,10 +244,40 @@ class LLMCallEvent:
upstream_started=kwargs.get("api_call_start_time") is not None,
provisional_span_name=f"{operation.value} {model}".strip(),
time_to_first_chunk_seconds=time_to_first_chunk_seconds(kwargs),
trace=caller_trace_controls(kwargs),
trace=trace,
session_id=caller_session_id(kwargs, trace),
)
def caller_session_id(kwargs: Mapping[str, object], trace: TraceControls) -> str | None:
"""The conversation id the caller sent (``litellm_session_id``, else the
``session_id`` trace control); ``None`` when the request carried none.
``get_litellm_params`` back-fills ``litellm_session_id`` from ``metadata.trace_id``
(which the proxy stamps with the OTel trace id) and ``missing_session_id: generate``
mints one into the body; neither is a caller conversation, so both are ignored,
while a ``langfuse_session_id`` header still counts under the generate policy.
``StandardLoggingPayload.session_id`` is never read: the payload drops the
generated marker, so a replayed minted id would pass for a caller's."""
params: Final[Mapping[str, object]] = as_str_mapping(kwargs.get("litellm_params")) or MappingProxyType({})
bodies: Final = tuple(
metadata
for key in ("metadata", "litellm_metadata")
if (metadata := as_str_mapping(params.get(key))) is not None
)
from_body: Final = tuple(session for body in bodies if (session := as_str(body.get("session_id"))))
minted: Final = frozenset(
session
for body in bodies
if body.get(SESSION_ID_GENERATED_METADATA_KEY) and (session := as_str(body.get("session_id")))
)
if minted:
return next((session for session in (trace.session_id, *from_body) if session and session not in minted), None)
explicit: Final = as_str(params.get("litellm_session_id"))
echoes_trace_id: Final = explicit is not None and any(as_str(body.get("trace_id")) == explicit for body in bodies)
return (None if echoes_trace_id else explicit) or trace.session_id or None
def time_to_first_chunk_seconds(kwargs: Mapping[str, Any]) -> float | None:
"""Seconds from the upstream request being issued (``api_call_start_time``)
to the first streamed chunk (``completion_start_time``); ``None`` for

View file

@ -407,6 +407,7 @@ class LLMCallSpanData:
call_type: str | None = None
request_route: str | None = None
trace: TraceControls = field(default_factory=TraceControls)
session_id: str | None = None
embedding_output: EmbeddingOutput | None = None
@classmethod
@ -417,6 +418,7 @@ class LLMCallSpanData:
time_to_first_chunk_seconds: float | None = None,
request_route: str | None = None,
trace: TraceControls | None = None,
session_id: str | None = None,
) -> LLMCallSpanData:
params: Final = cast(Mapping[str, object], payload.get("model_parameters") or {})
# The single parse of the request's metadata — the request-vs-provider
@ -463,6 +465,7 @@ class LLMCallSpanData:
call_type=call_type or None,
request_route=request_route or context.identity.request_route,
trace=trace or TraceControls(),
session_id=session_id or None,
embedding_output=embedding_output if capture_content else None,
)

View file

@ -67,7 +67,7 @@ class AmazonTwelveLabsPegasusConfig(AmazonInvokeConfig, BaseConfig):
optional_params["responseFormat"] = self._normalize_response_format(value)
return optional_params
def _normalize_response_format(self, value: Any) -> Any:
def _normalize_response_format(self, value: Any) -> object:
"""Normalize response_format to TwelveLabs format.
TwelveLabs expects:

View file

@ -358,13 +358,13 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
)
)
config_payload: Final[dict[str, Any]] = {
config_payload: Final[dict[str, object]] = {
"modules": modules if len(modules) > 1 else modules[0],
}
if stream_config:
config_payload["stream"] = stream_config
request_body: Final[dict[str, Any]] = {"config": config_payload}
request_body: Final[dict[str, object]] = {"config": config_payload}
if placeholder_values is not None:
request_body["placeholder_values"] = placeholder_values

View file

@ -83,8 +83,14 @@ def _oauth_token_error(code: str, status: int = 400) -> JSONResponse:
def _user_id_from_session_cookie(request: Request) -> str | None:
"""Return user_id from the UI ``token`` cookie (HS256-signed with
``master_key``), or None if missing/invalid.
"""Return user_id from the UI ``token`` cookie, or None if missing/invalid."""
user_id, _ = _session_identity_from_cookie(request)
return user_id
def _session_identity_from_cookie(request: Request) -> tuple[str | None, str | None]:
"""Return ``(user_id, session_key)`` from the UI ``token`` cookie
(HS256-signed with ``master_key``), or ``(None, None)`` if missing/invalid.
The /token endpoint in this file ALSO issues master-key-signed JWTs
(type="byok_session") for MCP-client-side use. They must not be
@ -98,10 +104,10 @@ def _user_id_from_session_cookie(request: Request) -> str | None:
from litellm.proxy.proxy_server import master_key
if not master_key:
return None
return None, None
token: Final = request.cookies.get("token")
if not token:
return None
return None, None
try:
payload: Final = jwt.decode(
token,
@ -113,21 +119,68 @@ def _user_id_from_session_cookie(request: Request) -> str | None:
options={"require": ["exp"]},
)
except jwt.InvalidTokenError:
return None
return None, None
if payload.get("type") == "byok_session":
return None
return None, None
if payload.get("login_method") not in ("sso", "username_password"):
return None
return None, None
user_id: Final = payload.get("user_id")
return user_id if isinstance(user_id, str) and user_id else None
if not isinstance(user_id, str) or not user_id:
return None, None
session_key: Final = payload.get("key")
return user_id, session_key if isinstance(session_key, str) and session_key else None
async def _session_key_is_live(session_key: str | None) -> bool:
"""Whether the session key embedded in the UI cookie still resolves.
The cookie JWT stays signature-valid until ``exp``; the DB-backed session
key inside it is what ``POST /session/logout`` and password-change
revocation actually kill. Trusting the signature alone would let a
logged-out cookie keep authorizing BYOK credential writes, so re-resolve
the key here.
EXPERIMENTAL_UI_LOGIN blob tokens (non-``sk-``) have no DB row and are
unrevocable by construction (scoped out of revocation); they pass through
on their bounded 10-minute lifetime, as before.
"""
from litellm.proxy._types import hash_token
from litellm.proxy.auth.auth_checks import get_key_object
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if session_key is None:
# Older cookies predating the ``key`` claim: nothing to resolve.
return True
if not session_key.startswith("sk-"):
return True
if prisma_client is None:
return True
try:
await get_key_object(
hashed_token=hash_token(session_key),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
except Exception:
return False
return True
async def _byok_session_auth(request: Request) -> UserAPIKeyAuth:
"""Require the UI session cookie. Programmatic BYOK management uses
"""Require the UI session cookie, with the embedded session key
re-resolved against the DB so a revoked (logged-out) session cannot
authorize BYOK writes. Programmatic BYOK management uses
``POST /v1/mcp/server/{id}/user-credential`` instead."""
user_id: Final = _user_id_from_session_cookie(request)
user_id, session_key = _session_identity_from_cookie(request)
if not user_id:
raise HTTPException(status_code=401, detail="login_required")
if not await _session_key_is_live(session_key):
raise HTTPException(status_code=401, detail="login_required")
return UserAPIKeyAuth(api_key="byok_session_cookie", user_id=user_id)

View file

@ -913,6 +913,7 @@ class LiteLLMRoutes(enum.Enum):
"/user/list", # org admins checked in endpoint; non-admins get 403
"/management/v1/users/bulk_delete", # proxy admins delete anyone, org admins only their orgs' users; others 403
"/user/password/change", # endpoint only ever writes the caller's own row
"/session/logout", # endpoint only ever revokes the caller's own session key
"/model/{model_id}/update",
"/prompt/list",
"/prompt/info",
@ -1948,6 +1949,10 @@ class ChangePasswordResponse(LiteLLMPydanticObjectBase):
message: str
class SessionLogoutResponse(LiteLLMPydanticObjectBase):
message: str
class DeleteUserRequest(LiteLLMPydanticObjectBase):
user_ids: list[str] # required

View file

@ -57,7 +57,7 @@ if TYPE_CHECKING:
from prisma import types as prisma_types
BREACH_RECHECK_INTERVAL: Final = timedelta(hours=24)
PASSWORD_RESET_ALLOWED_ROUTES: Final = ("/user/password/change",)
PASSWORD_RESET_ALLOWED_ROUTES: Final = ("/user/password/change", "/session/logout")
PASSWORD_SESSION_METADATA: Final = MappingProxyType({"login_method": "username_password"})

View file

@ -915,6 +915,10 @@ class RouteChecks:
if route == "/user/password/change":
return
# Self-service logout; the endpoint only revokes the caller's own session key.
if route == "/session/logout":
return
# Hard-block known write routes regardless of HTTP method (defensive
# — these are POSTs in practice, but pinning them here protects
# against future GET-shaped writes).

View file

@ -1,6 +1,6 @@
# litellm/proxy/guardrails/guardrail_hooks/pangea.py
import os
from typing import TYPE_CHECKING, Any, Final, cast
from typing import TYPE_CHECKING, Final
from fastapi import HTTPException
@ -230,7 +230,7 @@ class PangeaHandler(CustomGuardrail):
messages: Final = data.get("messages")
if messages is None:
return # No messages to check
input_messages = cast(list[dict[Any, Any]], messages)
input_messages = messages
else:
return

View file

@ -1583,6 +1583,23 @@ async def _update_single_user_helper(
response = inserted_user_row # pyright: ignore[reportAssignmentType] # insert_data returns a prisma row
if response is not None:
if "password" in non_default_values:
# An admin set this user's password, which implies the old one may be
# compromised; kill every existing UI session for the target. Revoke-all
# (no keep) — the caller is the admin, not the target, so the caller's
# own session is not among these.
from litellm.proxy.management_endpoints.session_endpoints import (
revoke_ui_session_keys,
)
target_user_id: Final = non_default_values.get("user_id")
if isinstance(target_user_id, str):
await revoke_ui_session_keys(
user_id=target_user_id,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
await _schedule_user_update_audit_log(
response=response,
existing_user_row=existing_user_row,

View file

@ -26,6 +26,7 @@ from litellm.proxy._types import (
from litellm.proxy.auth.login_utils import PASSWORD_SESSION_METADATA
from litellm.proxy.auth.password_policy import validate_password_not_breached, validate_password_policy
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.session_endpoints import revoke_ui_session_keys
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
from litellm.proxy.utils import hash_password, verify_password
from litellm.repositories.prisma_protocols import TableActions
@ -141,6 +142,15 @@ async def change_password(
}
await _user_table(prisma_client).update(where=find_user, data=password_update)
# The old password may have been compromised; revoke every other UI session
# so a holder of a stolen session token is cut off. The caller's own session
# is kept — they just proved they hold the current password.
await revoke_ui_session_keys(
user_id=user_id,
user_api_key_dict=user_api_key_dict,
keep_hashed_token=user_api_key_dict.token,
)
verbose_proxy_logger.info("Password changed via /user/password/change for user_id=%s", user_id)
await create_object_audit_log(
object_id=user_id,

View file

@ -0,0 +1,175 @@
"""
UI session revocation.
POST /session/logout — revoke the UI session key this request authenticated with.
revoke_ui_session_keys — revoke every UI session key a user holds (password writes).
Logging out of the dashboard was purely client-side (cookies cleared, redirect);
the DB-backed virtual key minted at login stayed valid until
LITELLM_UI_SESSION_DURATION elapsed, so a captured token kept working access
after logout, and changing a password did not invalidate existing sessions.
Deliberately NOT reusing /key/delete: its `can_modify_verification_token`
ownership checks can reject low-privilege roles, and a self-revoke endpoint
that takes no body cannot be aimed at other keys.
"""
from typing import TYPE_CHECKING, Annotated, Final, cast
from fastapi import APIRouter, Depends, HTTPException, Response
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
from litellm.proxy._types import (
CommonProxyErrors,
HTTPExceptionErrorDetail,
LiteLLM_VerificationToken,
SessionLogoutResponse,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import delete_cache_key_objects
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.key_management_endpoints import (
_persist_deleted_verification_tokens,
)
from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
if TYPE_CHECKING:
from prisma import types as prisma_types
router: Final = APIRouter()
_TOKEN_LIST: Final = TypeAdapter(list[str])
def _error_detail(message: str) -> HTTPExceptionErrorDetail:
detail: Final[HTTPExceptionErrorDetail] = {"error": message}
return detail
async def revoke_ui_session_keys(
user_id: str,
user_api_key_dict: UserAPIKeyAuth,
*,
keep_hashed_token: str | None = None,
litellm_changed_by: str | None = None,
) -> int:
"""Revoke every UI session key belonging to ``user_id``, except
``keep_hashed_token`` (the caller's own session on a self-service password
change; the other password-write paths revoke all).
Best-effort: the password write this runs after has already committed, so a
revocation failure is logged loudly rather than failing the request — the
unrevoked keys still expire at LITELLM_UI_SESSION_DURATION.
Returns the number of sessions revoked.
"""
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
if prisma_client is None:
return 0
try:
where_user_sessions: Final[prisma_types.LiteLLM_VerificationTokenWhereInput] = {
"user_id": user_id,
"team_id": UI_SESSION_TOKEN_TEAM_ID,
}
rows: Final = cast( # cast-ok: find_many returns prisma rows shaped like the pydantic model
"tuple[LiteLLM_VerificationToken, ...]",
tuple(await VerificationTokenRepository(prisma_client).table.find_many(where=where_user_sessions)),
)
revoked_rows: Final = tuple(row for row in rows if row.token is not None and row.token != keep_hashed_token)
if not revoked_rows:
return 0
revoked_tokens: Final = _TOKEN_LIST.validate_python(tuple(row.token for row in revoked_rows))
await _persist_deleted_verification_tokens(
keys=revoked_rows,
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
where_revoked: Final[prisma_types.LiteLLM_VerificationTokenWhereInput] = {"token": {"in": revoked_tokens}}
await VerificationTokenRepository(prisma_client).table.delete_many(where=where_revoked)
await delete_cache_key_objects(
hashed_tokens=revoked_tokens,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
verbose_proxy_logger.info(
"Revoked %s UI session key(s) for user_id=%s after password change",
len(revoked_tokens),
user_id,
)
return len(revoked_tokens)
except Exception: # noqa: BLE001 # the password write committed; revocation must not undo that
verbose_proxy_logger.exception(
"Failed to revoke UI session keys for user_id=%s; existing sessions remain valid until they expire",
user_id,
)
return 0
@router.post(
"/session/logout",
tags=("UI Session",),
)
async def session_logout(
response: Response,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> SessionLogoutResponse:
"""
Revoke the UI session key this request authenticated with.
Only accepts UI session keys (minted by dashboard login); any other
credential is refused, so this can never be used to delete arbitrary keys.
Revokes only the presented session, not the user's other sessions.
Idempotent: logging out an already-revoked session succeeds.
"""
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
if prisma_client is None:
raise HTTPException(
status_code=500,
detail=_error_detail(CommonProxyErrors.db_not_connected_error.value),
)
if user_api_key_dict.team_id != UI_SESSION_TOKEN_TEAM_ID:
raise HTTPException(
status_code=403,
detail=_error_detail("Only UI session tokens can be revoked through this endpoint."),
)
hashed_token: Final = user_api_key_dict.token
revoked = False
if hashed_token is not None:
where_token: Final[prisma_types.LiteLLM_VerificationTokenWhereUniqueInput] = {"token": hashed_token}
row: Final = await VerificationTokenRepository(prisma_client).table.find_unique(where=where_token)
# A missing row means the session is already revoked (or an
# EXPERIMENTAL_UI_LOGIN blob token); logout is idempotent either way.
if row is not None:
caller_row: Final = cast( # cast-ok: find_unique returns a prisma row shaped like the pydantic model
"LiteLLM_VerificationToken", row
)
await _persist_deleted_verification_tokens(
keys=(caller_row,),
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
)
await VerificationTokenRepository(prisma_client).table.delete_many(where=where_token)
revoked = True
await delete_cache_key_objects(
hashed_tokens=(hashed_token,),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# The server set this cookie at login (set_session_token_cookie); clear it
# here too so logout works even if the client-side clear is skipped.
response.delete_cookie("token")
return SessionLogoutResponse(
message="Session revoked." if revoked else "Session already revoked.",
)

View file

@ -630,6 +630,9 @@ from litellm.proxy.management_endpoints.prompt_caching_requests import (
from litellm.proxy.management_endpoints.router_settings_endpoints import (
router as router_settings_router,
)
from litellm.proxy.management_endpoints.session_endpoints import (
router as session_management_router,
)
from litellm.proxy.management_endpoints.tag_management_endpoints import (
router as tag_management_router,
)
@ -17012,6 +17015,19 @@ async def claim_onboarding_link(data: InvitationClaim, request: Request):
if user_obj and hasattr(user_obj, "__dict__"):
user_obj.__dict__.pop("password", None)
# The password just changed via an invitation/reset link; any UI session
# minted under the old password may be in hostile hands. Revoke them all —
# the caller holds only the short-lived onboarding JWT, and the fresh
# session key is minted below, after this sweep.
from litellm.proxy.management_endpoints.session_endpoints import (
revoke_ui_session_keys,
)
await revoke_ui_session_keys(
user_id=invite_obj.user_id,
user_api_key_dict=UserAPIKeyAuth(user_id=invite_obj.user_id),
)
try:
jwt_token: Final = await _generate_onboarding_ui_session_token(user_obj=user_obj)
except Exception as e:
@ -19431,6 +19447,7 @@ app.include_router(health_router)
app.include_router(key_management_router)
app.include_router(internal_user_router)
app.include_router(password_management_router)
app.include_router(session_management_router)
app.include_router(team_router)
app.include_router(ui_sso_router)
app.include_router(organization_router)

View file

@ -52,7 +52,7 @@ def get_instance_fn(value: str, config_file_path: str | None = None) -> Any:
module = importlib.import_module(module_name)
# Get the instance from the module
instance: Final = getattr(module, instance_name)
instance: Final[object] = getattr(module, instance_name)
return instance
except ImportError as e:
@ -167,7 +167,7 @@ def _load_instance_from_remote_storage(remote_url: str, config_file_path: str |
spec.loader.exec_module(module)
# Get the instance
instance: Final = getattr(module, instance_name)
instance: Final[object] = getattr(module, instance_name)
# Clean up the temporary file
try:

View file

@ -2041,6 +2041,75 @@
"tests/integration/observability/test_callback_delivery.py::test_response_survives_raising_success_deployment_hook[videos]": [
"other.observability.callbacks.raising_success_deployment_hook_keeps_response"
],
"tests/integration/observability/test_otel_conversation_id.py::test_chat_completion_sdk_body_litellm_session_id_lands_as_conversation_id": [
"other.observability.otel.conversation_id_from_body_session_id_chat_sdk"
],
"tests/integration/observability/test_otel_conversation_id.py::test_chat_stream_async_sdk_x_litellm_session_id_header_lands_as_conversation_id": [
"other.observability.otel.conversation_id_from_header_chat_stream_async_sdk"
],
"tests/integration/observability/test_otel_conversation_id.py::test_messages_sdk_x_litellm_session_id_header_lands_as_conversation_id": [
"other.observability.otel.conversation_id_from_header_messages_sdk"
],
"tests/integration/observability/test_otel_conversation_id.py::test_messages_stream_async_sdk_langfuse_session_id_header_lands_as_conversation_id": [
"other.observability.otel.conversation_id_from_langfuse_header_messages_stream_async_sdk"
],
"tests/integration/observability/test_otel_conversation_id.py::test_responses_sdk_x_litellm_session_id_header_lands_as_conversation_id": [
"other.observability.otel.conversation_id_from_header_responses_sdk"
],
"tests/integration/observability/test_otel_conversation_id.py::test_responses_stream_raw_metadata_session_id_lands_as_conversation_id": [
"other.observability.otel.conversation_id_from_metadata_responses_stream_raw"
],
"tests/integration/observability/test_otel_conversation_id.py::test_chat_raw_metadata_session_id_lands_as_conversation_id": [
"other.observability.otel.conversation_id_from_metadata_chat_raw"
],
"tests/integration/observability/test_otel_conversation_id.py::test_integer_and_list_litellm_session_id_match_the_spend_row_or_are_dropped_together": [
"other.observability.otel.conversation_id_non_string_session_ids_match_spend_row"
],
"tests/integration/observability/test_otel_conversation_id.py::test_empty_string_litellm_session_id_leaves_the_span_without_a_conversation_id": [
"other.observability.otel.conversation_id_empty_string_session_id_is_omitted"
],
"tests/integration/observability/test_otel_conversation_id.py::test_five_kilobyte_session_header_round_trips_to_the_span_and_the_spend_row": [
"other.observability.otel.conversation_id_five_kilobyte_header_round_trips"
],
"tests/integration/observability/test_otel_conversation_id.py::test_duplicate_session_header_lands_once_and_unchanged": [
"other.observability.otel.conversation_id_duplicate_header_lands_once"
],
"tests/integration/observability/test_otel_conversation_id.py::test_unauthenticated_request_with_session_header_is_rejected_and_leaves_no_span": [
"other.observability.otel.conversation_id_unauthenticated_request_leaves_no_span"
],
"tests/integration/observability/test_otel_conversation_id.py::test_sink_rejecting_with_403_drops_those_spans_and_later_spans_still_land": [
"other.observability.otel.conversation_id_survives_sink_rejection"
],
"tests/integration/observability/test_otel_conversation_id.py::test_request_without_any_session_input_has_no_conversation_id": [
"other.observability.otel.conversation_id_absent_without_caller_session"
],
"tests/integration/observability/test_otel_conversation_id.py::test_generate_policy_minted_session_id_reaches_the_spend_row_but_not_the_span": [
"other.observability.otel.conversation_id_ignores_generated_session_id"
],
"tests/integration/observability/test_otel_conversation_id.py::test_generate_policy_keeps_the_langfuse_session_header_as_conversation_id": [
"other.observability.otel.conversation_id_langfuse_header_wins_over_generated"
],
"tests/integration/observability/test_otel_conversation_id.py::test_header_body_and_metadata_session_ids_resolve_to_the_same_id_as_the_spend_row": [
"other.observability.otel.conversation_id_header_precedence_matches_spend_row"
],
"tests/integration/observability/test_otel_conversation_id.py::test_three_identical_requests_produce_one_span_each_with_the_same_conversation_id": [
"other.observability.otel.conversation_id_repeated_requests_log_once_each"
],
"tests/integration/observability/test_otel_conversation_id.py::test_metadata_trace_id_alone_fills_the_spend_row_but_not_the_span": [
"other.observability.otel.conversation_id_ignores_trace_id_backfill"
],
"tests/integration/observability/test_otel_conversation_id.py::test_sink_outage_during_a_mixed_burst_lands_every_response_exactly_once_after_recovery": [
"other.observability.otel.conversation_id_sink_outage_recovers_exactly_once"
],
"tests/integration/observability/test_otel_conversation_id.py::test_slow_sink_during_a_burst_lands_every_response_exactly_once": [
"other.observability.otel.conversation_id_slow_sink_no_duplicates"
],
"tests/integration/observability/test_otel_conversation_id.py::test_killing_one_of_two_workers_mid_burst_keeps_serving_and_never_duplicates_a_span": [
"other.observability.otel.conversation_id_survives_worker_kill"
],
"tests/integration/observability/test_otel_conversation_id.py::test_terminating_the_proxy_right_after_a_burst_flushes_every_span_before_exit": [
"other.observability.otel.conversation_id_flushes_on_shutdown"
],
"tests/integration/management/test_user_updates_wedged_coordination_redis.py::test_user_budget_updates_return_promptly_while_coordination_redis_is_wedged": [
"mgmt.user.update.budget_change_returns_promptly_with_wedged_coordination_redis",
"mgmt.user.bulk_update.budget_change_returns_promptly_with_wedged_coordination_redis",

View file

@ -0,0 +1,830 @@
import asyncio
import base64
import json
import os
import re
import signal
import threading
import uuid
from collections import deque
from collections.abc import Iterator, Mapping, Sequence
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import anthropic
import httpx
import openai
import psutil
import pytest
import yaml
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.database import read_rows
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
MARKER: Final = re.compile(rb"otelconv-[0-9a-f]{32}")
CONVERSATION: Final = "gen_ai.conversation.id"
def _marker() -> str:
return "otelconv-" + uuid.uuid4().hex
def _chat_reply(identity: str, stream: bool) -> Reply:
if not stream:
return Reply(
body=json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "conversation ok"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9},
}
).encode()
)
chunk: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"}
return Reply(
content_type="text/event-stream",
chunks=(
b"data: "
+ json.dumps(
{**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "conversation"}}]}
).encode()
+ b"\n\n",
b"data: "
+ json.dumps(
{
**chunk,
"choices": [{"index": 0, "delta": {"content": " ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9},
}
).encode()
+ b"\n\n",
b"data: [DONE]\n\n",
),
)
def _responses_reply(identity: str, stream: bool) -> Reply:
response: Final = {
"id": identity,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"id": "msg_" + identity,
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "conversation ok", "annotations": []}],
}
],
"usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9},
}
if not stream:
return Reply(body=json.dumps(response).encode())
events: Final = (
{
"type": "response.created",
"sequence_number": 0,
"response": {**response, "status": "in_progress", "output": []},
},
{
"type": "response.output_text.delta",
"sequence_number": 1,
"item_id": "msg_" + identity,
"output_index": 0,
"content_index": 0,
"delta": "conversation ok",
},
{"type": "response.completed", "sequence_number": 2, "response": response},
)
return Reply(
content_type="text/event-stream",
chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events),
)
def _decoded_responses_id(identity: str) -> str:
try:
return base64.b64decode(identity.removeprefix("resp_").encode()).decode()
except (ValueError, UnicodeDecodeError):
return identity
def _canonical_id(identity: str) -> str:
return _decoded_responses_id(identity).rpartition("response_id:")[2]
def _sse_events(text: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(
json.loads(line[6:]) for line in text.splitlines() if line.startswith("data: ") and line != "data: [DONE]"
)
def _upstream(request: Request) -> Reply:
found: Final = MARKER.search(request.body)
if found is None:
return Reply(status=404, body=b'{"error":"no marker"}')
marker: Final = found.group(0).decode()
stream: Final = json.loads(request.body).get("stream") is True
if request.target.endswith("/responses"):
return _responses_reply(f"resp_{marker}", stream)
return _chat_reply(f"chatcmpl-{marker}", stream)
@dataclass(frozen=True, slots=True)
class Collector:
wire: Wire
outage: threading.Event
rejection: threading.Event
slow: threading.Event
accepted: Sequence[Request]
guard: threading.Lock
def attributes(self) -> tuple[dict[str, dict[str, JsonValue]], ...]:
with self.guard:
batches: Final = tuple(self.accepted)
return tuple(
{attribute["key"]: attribute["value"] for attribute in span.get("attributes", ())}
for batch in batches
for resource in json.loads(batch.body)["resourceSpans"]
for scope in resource["scopeSpans"]
for span in scope["spans"]
)
def spans(self, response_id: str) -> tuple[dict[str, dict[str, JsonValue]], ...]:
return tuple(
attributes
for attributes in self.attributes()
if isinstance(logged := attributes.get("gen_ai.response.id", {}).get("stringValue"), str)
and _canonical_id(logged) == _canonical_id(response_id)
)
def conversation_ids(self, response_id: str) -> tuple[str | None, ...]:
return tuple(
attributes[CONVERSATION]["stringValue"] if CONVERSATION in attributes else None
for attributes in self.spans(response_id)
)
def single_span(self, response_id: str) -> str | None:
return eventually(lambda: self.conversation_ids(response_id), lambda values: len(values) == 1, seconds=30)[0]
def logged_id(self, response_id: str) -> str:
spans: Final = eventually(lambda: self.spans(response_id), lambda values: len(values) == 1, seconds=30)
return str(spans[0]["gen_ai.response.id"]["stringValue"])
@pytest.fixture(scope="session")
def collector() -> Iterator[Collector]:
outage: Final = threading.Event()
rejection: Final = threading.Event()
slow: Final = threading.Event()
accepted: Final[deque[Request]] = deque() # mutable-ok: sink thread appends each accepted batch
guard: Final = threading.Lock()
def sink(request: Request) -> Reply:
if slow.is_set():
threading.Event().wait(1.5)
if outage.is_set():
return Reply(status=503, body=b'{"error":"sink down"}')
if rejection.is_set():
return Reply(status=403, body=b'{"error":"forbidden"}')
with guard:
accepted.append(request)
return Reply()
with wire_server(sink) as wire:
yield Collector(wire, outage, rejection, slow, accepted, guard)
@pytest.fixture(scope="session")
def provider() -> Iterator[Wire]:
with wire_server(_upstream) as wire:
yield wire
@dataclass(frozen=True, slots=True)
class Rig:
proxy: Gateway
process: OwnedProxy
model: str
upstream: Wire
sink: Collector
def openai_client(self) -> openai.OpenAI:
return openai.OpenAI(base_url=str(self.proxy.client.base_url) + "/v1", api_key=self.proxy.key, max_retries=0)
def async_openai_client(self) -> openai.AsyncOpenAI:
return openai.AsyncOpenAI(
base_url=str(self.proxy.client.base_url) + "/v1", api_key=self.proxy.key, max_retries=0
)
def anthropic_client(self) -> anthropic.Anthropic:
return anthropic.Anthropic(base_url=str(self.proxy.client.base_url), api_key=self.proxy.key, max_retries=0)
def async_anthropic_client(self) -> anthropic.AsyncAnthropic:
return anthropic.AsyncAnthropic(base_url=str(self.proxy.client.base_url), api_key=self.proxy.key, max_retries=0)
def chat(
self,
marker: str,
*,
headers: Mapping[str, str] | None = None,
key: str | None = None,
**extra: JsonValue,
) -> httpx.Response:
return self.proxy.request(
"POST",
"/v1/chat/completions",
{
"model": self.model,
"messages": [{"role": "user", "content": marker}],
"cache": {"no-cache": True},
**extra,
},
headers=headers,
key=key,
)
def upstream_bodies(self, marker: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(json.loads(request.body) for request in self.upstream.drain() if marker.encode() in request.body)
def spend_session(self, response_id: str) -> str | None:
rows: Final = eventually(
lambda: read_rows('SELECT session_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (response_id,)),
lambda values: len(values) == 1,
seconds=70,
)
value: Final = rows[0]["session_id"]
assert value is None or isinstance(value, str), rows
return value
def spend_request_ids(self, session: str) -> tuple[str, ...]:
rows: Final = read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE session_id=%s', (session,))
return tuple(str(row["request_id"]) for row in rows)
@dataclass(frozen=True, slots=True)
class RigFactory:
provider: Wire
sink: Collector
directory: Path
settings: Mapping[str, JsonValue]
workers: int
def start(self) -> Iterator[Rig]:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update({"callbacks": ["otel"]})
config["general_settings"].update({"disable_model_info_refresh": True, **self.settings})
config["callback_settings"] = {
"otel": {"exporter": "http/json", "endpoint": self.sink.wire.url, "mapper_names": ["genai"]},
}
path: Final = self.directory / f"otel-{uuid.uuid4().hex}.yaml"
path.write_text(yaml.safe_dump(config))
overrides: Final = {"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300"}
with (
gateway_from_environment() as gateway,
owned_proxy_process(gateway, self.directory, overrides, config=path, workers=self.workers) as owned,
owned.gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=self.provider.url + "/v1")
yield Rig(owned.gateway, owned, model, self.provider, self.sink)
@pytest.fixture(scope="session")
def rig(provider: Wire, collector: Collector, tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
yield from RigFactory(provider, collector, tmp_path_factory.mktemp("otel"), {}, 2).start()
@pytest.fixture(scope="session")
def generating_rig(provider: Wire, collector: Collector, tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
factory: Final = RigFactory(
provider, collector, tmp_path_factory.mktemp("otel-generate"), {"missing_session_id": "generate"}, 2
)
yield from factory.start()
@pytest.fixture(scope="session")
def two_worker_rig(provider: Wire, collector: Collector, tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
yield from RigFactory(provider, collector, tmp_path_factory.mktemp("otel-workers"), {}, 2).start()
def _assert_upstream_clean(rig: Rig, marker: str, session: str) -> None:
bodies: Final = rig.upstream_bodies(marker)
assert len(bodies) == 1, bodies
assert session not in json.dumps(bodies[0]), bodies[0]
@pytest.mark.covers("other.observability.otel.conversation_id_from_body_session_id_chat_sdk")
def test_chat_completion_sdk_body_litellm_session_id_lands_as_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
completion: Final = rig.openai_client().chat.completions.create(
model=rig.model,
messages=[{"role": "user", "content": marker}],
extra_body={"litellm_session_id": session, "cache": {"no-cache": True}},
)
assert completion.id == f"chatcmpl-{marker}", completion
assert completion.choices[0].message.content == "conversation ok", completion
assert rig.sink.single_span(completion.id) == session
assert rig.spend_session(completion.id) == session
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_from_header_chat_stream_async_sdk")
def test_chat_stream_async_sdk_x_litellm_session_id_header_lands_as_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
async def consume() -> tuple[str, str]:
stream: Final = await rig.async_openai_client().chat.completions.create(
model=rig.model,
messages=[{"role": "user", "content": marker}],
stream=True,
extra_headers={"x-litellm-session-id": session},
extra_body={"cache": {"no-cache": True}},
)
chunks: Final = [chunk async for chunk in stream]
return chunks[0].id, "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
identity, text = asyncio.run(consume())
assert identity == f"chatcmpl-{marker}", identity
assert text == "conversation ok", text
assert rig.sink.single_span(identity) == session
assert rig.spend_session(identity) == session
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_from_header_messages_sdk")
def test_messages_sdk_x_litellm_session_id_header_lands_as_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
message: Final = rig.anthropic_client().messages.create(
model=rig.model,
max_tokens=16,
messages=[{"role": "user", "content": marker}],
extra_headers={"x-litellm-session-id": session},
)
assert message.content[0].type == "text" and message.content[0].text == "conversation ok", message
assert rig.sink.single_span(message.id) == session
assert rig.spend_session(message.id) == session
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_from_langfuse_header_messages_stream_async_sdk")
def test_messages_stream_async_sdk_langfuse_session_id_header_lands_as_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
async def consume() -> tuple[str, str]:
async with rig.async_anthropic_client().messages.stream(
model=rig.model,
max_tokens=16,
messages=[{"role": "user", "content": marker}],
extra_headers={"langfuse_session_id": session},
) as stream:
text: Final = "".join([piece async for piece in stream.text_stream])
return (await stream.get_final_message()).id, text
identity, text = asyncio.run(consume())
assert text == "conversation ok", text
assert rig.sink.single_span(identity) == session
assert rig.spend_session(identity), identity
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_from_header_responses_sdk")
def test_responses_sdk_x_litellm_session_id_header_lands_as_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
response: Final = rig.openai_client().responses.create(
model=rig.model, input=marker, extra_headers={"x-litellm-session-id": session}
)
assert response.output[0].id == f"msg_resp_{marker}", response
assert response.output_text == "conversation ok", response
assert rig.sink.single_span(response.id) == session
assert rig.spend_session(response.id) == session
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_from_metadata_responses_stream_raw")
def test_responses_stream_raw_metadata_session_id_lands_as_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
with rig.proxy.client.stream(
"POST",
"/v1/responses",
json={"model": rig.model, "input": marker, "stream": True, "metadata": {"session_id": session}},
headers={"Authorization": f"Bearer {rig.proxy.key}"},
) as response:
body: Final = response.read().decode()
assert response.status_code == 200, body
events: Final = _sse_events(body)
completed: Final = tuple(event for event in events if event["type"] == "response.completed")
assert len(completed) == 1, events
assert completed[0]["response"]["output"][0]["id"] == f"msg_resp_{marker}", completed
assert str(completed[0]["response"]["id"]).startswith("resp_"), completed
assert rig.upstream_bodies(marker) == (
{"model": "gpt-4o-mini", "input": marker, "metadata": {"session_id": session}, "stream": True},
)
assert rig.sink.single_span(f"resp_{marker}") == session
assert rig.spend_session(rig.sink.logged_id(f"resp_{marker}")) == session
@pytest.mark.covers("other.observability.otel.conversation_id_from_metadata_chat_raw")
def test_chat_raw_metadata_session_id_lands_as_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
response: Final = rig.chat(marker, metadata={"session_id": session})
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert identity == f"chatcmpl-{marker}", response.text
assert rig.sink.single_span(identity) == session
assert rig.spend_session(identity) == session
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_non_string_session_ids_match_spend_row")
def test_integer_and_list_litellm_session_id_match_the_spend_row_or_are_dropped_together(rig: Rig) -> None:
for odd in (123, ["a", "b"]):
marker: Final = _marker()
response: Final = rig.chat(marker, litellm_session_id=odd)
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert rig.sink.single_span(identity) == rig.spend_session(identity), (odd, rig.sink.conversation_ids(identity))
assert len(rig.upstream_bodies(marker)) == 1
@pytest.mark.covers("other.observability.otel.conversation_id_empty_string_session_id_is_omitted")
def test_empty_string_litellm_session_id_leaves_the_span_without_a_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
response: Final = rig.chat(marker, litellm_session_id="")
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert rig.sink.single_span(identity) is None
assert rig.spend_session(identity), response.text
@pytest.mark.covers("other.observability.otel.conversation_id_five_kilobyte_header_round_trips")
def test_five_kilobyte_session_header_round_trips_to_the_span_and_the_spend_row(rig: Rig) -> None:
marker: Final = _marker()
session: Final = ("s" * 5000) + uuid.uuid4().hex
response: Final = rig.chat(marker, headers={"x-litellm-session-id": session})
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert rig.sink.single_span(identity) == session
assert rig.spend_session(identity) == session
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_duplicate_header_lands_once")
def test_duplicate_session_header_lands_once_and_unchanged(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
response: Final = rig.proxy.client.post(
"/v1/chat/completions",
json={"model": rig.model, "messages": [{"role": "user", "content": marker}], "cache": {"no-cache": True}},
headers=[
("Authorization", f"Bearer {rig.proxy.key}"),
("x-litellm-session-id", session),
("x-litellm-session-id", session),
],
)
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert rig.sink.single_span(identity) == session
assert rig.spend_session(identity) == session
_assert_upstream_clean(rig, marker, session)
@pytest.mark.covers("other.observability.otel.conversation_id_unauthenticated_request_leaves_no_span")
def test_unauthenticated_request_with_session_header_is_rejected_and_leaves_no_span(rig: Rig) -> None:
marker: Final = _marker()
response: Final = rig.chat(marker, headers={"x-litellm-session-id": "conv-" + uuid.uuid4().hex}, key="sk-wrong")
assert response.status_code == 401, response.text
assert rig.upstream_bodies(marker) == ()
control: Final = rig.chat(marker)
assert control.status_code == 200, control.text
assert rig.sink.single_span(control.json()["id"]) is None
assert rig.spend_session(control.json()["id"]), control.text
@pytest.mark.covers("other.observability.otel.conversation_id_survives_sink_rejection")
def test_sink_rejecting_with_403_drops_those_spans_and_later_spans_still_land(rig: Rig) -> None:
rig.sink.rejection.set()
try:
rejected: Final = rig.chat(_marker(), headers={"x-litellm-session-id": "conv-rejected"})
assert rejected.status_code == 200, rejected.text
eventually(lambda: any(request.body for request in rig.sink.wire.drain()), lambda seen: seen, seconds=30)
finally:
rig.sink.rejection.clear()
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
response: Final = rig.chat(marker, headers={"x-litellm-session-id": session})
assert response.status_code == 200, response.text
assert rig.sink.single_span(response.json()["id"]) == session
@pytest.mark.covers("other.observability.otel.conversation_id_absent_without_caller_session")
def test_request_without_any_session_input_has_no_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
response: Final = rig.chat(marker)
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert rig.sink.single_span(identity) is None
assert rig.spend_session(identity), response.text
@pytest.mark.covers("other.observability.otel.conversation_id_ignores_generated_session_id")
def test_generate_policy_minted_session_id_reaches_the_spend_row_but_not_the_span(generating_rig: Rig) -> None:
marker: Final = _marker()
response: Final = generating_rig.chat(marker)
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
minted: Final = generating_rig.spend_session(identity)
assert minted, response.text
assert generating_rig.sink.single_span(identity) is None
@pytest.mark.covers("other.observability.otel.conversation_id_langfuse_header_wins_over_generated")
def test_generate_policy_keeps_the_langfuse_session_header_as_conversation_id(generating_rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
response: Final = generating_rig.chat(marker, headers={"langfuse_session_id": session})
assert response.status_code == 200, response.text
assert generating_rig.sink.single_span(response.json()["id"]) == session
@pytest.mark.covers("other.observability.otel.conversation_id_header_precedence_matches_spend_row")
def test_header_body_and_metadata_session_ids_resolve_to_the_same_id_as_the_spend_row(rig: Rig) -> None:
marker: Final = _marker()
header: Final = "conv-header-" + uuid.uuid4().hex
response: Final = rig.chat(
marker,
headers={"x-litellm-session-id": header},
litellm_session_id="conv-body-" + uuid.uuid4().hex,
metadata={"session_id": "conv-meta-" + uuid.uuid4().hex},
)
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert rig.sink.single_span(identity) == header
assert rig.spend_session(identity) == header
@pytest.mark.covers("other.observability.otel.conversation_id_repeated_requests_log_once_each")
def test_three_identical_requests_produce_one_span_each_with_the_same_conversation_id(rig: Rig) -> None:
marker: Final = _marker()
session: Final = "conv-" + uuid.uuid4().hex
responses: Final = tuple(rig.chat(marker, headers={"x-litellm-session-id": session}) for _ in range(3))
assert all(response.status_code == 200 for response in responses), [response.text for response in responses]
identity: Final = f"chatcmpl-{marker}"
spans: Final = eventually(lambda: rig.sink.conversation_ids(identity), lambda values: len(values) == 3, seconds=30)
assert spans == (session, session, session), spans
assert len(rig.upstream_bodies(marker)) == 3
@pytest.mark.covers("other.observability.otel.conversation_id_ignores_trace_id_backfill")
def test_metadata_trace_id_alone_fills_the_spend_row_but_not_the_span(rig: Rig) -> None:
marker: Final = _marker()
trace: Final = "trace-" + uuid.uuid4().hex
response: Final = rig.chat(marker, metadata={"trace_id": trace})
assert response.status_code == 200, response.text
identity: Final = response.json()["id"]
assert rig.sink.single_span(identity) is None
assert rig.spend_session(identity) == trace
def _chat_id(response: httpx.Response) -> str:
if not response.headers.get("content-type", "").startswith("text/event-stream"):
return response.json()["id"]
identities: Final = frozenset(str(event["id"]) for event in _sse_events(response.text))
assert len(identities) == 1, response.text
return next(iter(identities))
def _responses_id(response: httpx.Response) -> str:
if not response.headers.get("content-type", "").startswith("text/event-stream"):
return response.json()["id"]
completed: Final = tuple(
event["response"]["id"] for event in _sse_events(response.text) if event.get("type") == "response.completed"
)
assert len(completed) == 1, response.text
return str(completed[0])
def _message_id(response: httpx.Response) -> str:
if not response.headers.get("content-type", "").startswith("text/event-stream"):
return response.json()["id"]
starts: Final = tuple(
event["message"]["id"] for event in _sse_events(response.text) if event.get("type") == "message_start"
)
assert len(starts) == 1, response.text
return starts[0]
def _burst(rig: Rig, count: int, session_for: Mapping[int, str]) -> tuple[tuple[int, str, str | None], ...]:
markers: Final = tuple(_marker() for _ in range(count))
def one(index: int) -> tuple[int, str, str | None]:
marker: Final = markers[index]
headers: Final = {"Authorization": f"Bearer {rig.proxy.key}", "x-litellm-session-id": session_for[index]}
route: Final = index % 3
try:
if route == 0:
response: Final = rig.proxy.client.post(
"/v1/chat/completions",
json={
"model": rig.model,
"messages": [{"role": "user", "content": marker}],
"stream": index % 2 == 0,
},
headers=headers,
)
response.read()
if response.status_code != 200:
return index, marker, response.text
return index, _chat_id(response), None
if route == 1:
response = rig.proxy.client.post(
"/v1/responses",
json={"model": rig.model, "input": marker, "stream": index % 2 == 0},
headers=headers,
)
response.read()
if response.status_code != 200:
return index, marker, response.text
return index, _responses_id(response), None
response = rig.proxy.client.post(
"/v1/messages",
json={
"model": rig.model,
"max_tokens": 16,
"messages": [{"role": "user", "content": marker}],
"stream": index % 2 == 0,
},
headers=headers,
)
response.read()
if response.status_code != 200:
return index, marker, response.text
return index, _message_id(response), None
except httpx.HTTPError as error:
return index, marker, repr(error)
with ThreadPoolExecutor(max_workers=10) as pool:
return tuple(pool.map(one, range(count)))
def _is_encrypted_responses_id(identity: str) -> bool:
return identity.startswith("resp_") and _decoded_responses_id(identity) == identity
def _landed(rig: Rig, expected: Mapping[str, str]) -> dict[str, tuple[str, ...]]:
spans: Final = rig.sink.attributes()
return {
session: tuple(
_canonical_id(str(attributes["gen_ai.response.id"]["stringValue"]))
for attributes in spans
if attributes.get(CONVERSATION, {}).get("stringValue") == session and "gen_ai.response.id" in attributes
)
for session in expected.values()
}
def _assert_exactly_once(rig: Rig, expected: Mapping[str, str], landed: Mapping[str, tuple[str, ...]]) -> None:
spend: Final = eventually(
lambda: {
session: tuple(_canonical_id(identity) for identity in rig.spend_request_ids(session))
for session in expected.values()
},
lambda rows: all(len(values) >= 1 for values in rows.values()),
seconds=70,
)
assert landed == spend, (landed, spend)
assert all(len(values) == 1 for values in landed.values()), landed
caller_visible: Final = {
session: (_canonical_id(identity),)
for identity, session in expected.items()
if not _is_encrypted_responses_id(identity)
}
assert {session: landed[session] for session in caller_visible} == caller_visible, landed
@pytest.mark.covers("other.observability.otel.conversation_id_sink_outage_recovers_exactly_once")
def test_sink_outage_during_a_mixed_burst_lands_every_response_exactly_once_after_recovery(rig: Rig) -> None:
sessions: Final = {index: f"conv-{index}-{uuid.uuid4().hex}" for index in range(30)}
rig.sink.outage.set()
try:
health_down: Final = rig.proxy.request("GET", "/health/services", params={"service": "otel"})
results: Final = _burst(rig, 30, sessions)
assert all(error is None for _, _, error in results), [error for _, _, error in results if error]
eventually(lambda: any(True for _ in rig.sink.wire.drain()), lambda seen: seen, seconds=30)
finally:
rig.sink.outage.clear()
assert health_down.status_code == 200, health_down.text
expected: Final = {identity: sessions[index] for index, identity, _ in results}
landed: Final = eventually(
lambda: _landed(rig, expected), lambda seen: all(len(values) >= 1 for values in seen.values()), seconds=80
)
_assert_exactly_once(rig, expected, landed)
@pytest.mark.covers("other.observability.otel.conversation_id_slow_sink_no_duplicates")
def test_slow_sink_during_a_burst_lands_every_response_exactly_once(rig: Rig) -> None:
sessions: Final = {index: f"conv-{index}-{uuid.uuid4().hex}" for index in range(20)}
rig.sink.slow.set()
try:
results: Final = _burst(rig, 20, sessions)
assert all(error is None for _, _, error in results), [error for _, _, error in results if error]
expected: Final = {identity: sessions[index] for index, identity, _ in results}
landed: Final = eventually(
lambda: _landed(rig, expected), lambda seen: all(len(values) >= 1 for values in seen.values()), seconds=80
)
finally:
rig.sink.slow.clear()
_assert_exactly_once(rig, expected, landed)
@pytest.mark.covers("other.observability.otel.conversation_id_survives_worker_kill")
def test_killing_one_of_two_workers_mid_burst_keeps_serving_and_never_duplicates_a_span(two_worker_rig: Rig) -> None:
rig: Final = two_worker_rig
root: Final = psutil.Process(rig.process.process.pid)
workers: Final = eventually(
lambda: tuple(child for child in root.children() if "resource_tracker" not in " ".join(child.cmdline())),
lambda found: len(found) == 2,
seconds=30,
)
sessions: Final = {index: f"conv-{index}-{uuid.uuid4().hex}" for index in range(24)}
markers: Final = tuple(_marker() for _ in range(24))
def one(index: int) -> tuple[str, str | None]:
if index == 8:
os.kill(workers[0].pid, signal.SIGKILL)
try:
response: Final = rig.chat(markers[index], headers={"x-litellm-session-id": sessions[index]})
return f"chatcmpl-{markers[index]}", None if response.status_code == 200 else response.text
except httpx.HTTPError as error:
return f"chatcmpl-{markers[index]}", repr(error)
with ThreadPoolExecutor(max_workers=6) as pool:
results: Final = tuple(pool.map(one, range(24)))
assert rig.process.process.poll() is None, "Proxy root exited after a worker was killed"
after: Final = rig.chat(_marker(), headers={"x-litellm-session-id": "conv-after-kill"})
assert after.status_code == 200, after.text
assert rig.sink.single_span(after.json()["id"]) == "conv-after-kill"
failures: Final = tuple(error for _, error in results if error)
assert all(error.startswith(("ReadError(", "RemoteProtocolError(", "ConnectError(")) for error in failures), (
failures
)
assert len(failures) <= 6, failures
served: Final = {identity: sessions[index] for index, (identity, error) in enumerate(results) if error is None}
assert len(served) >= 18, results
settled: Final = {
identity: sessions[index] for index, (identity, error) in enumerate(results) if index > 14 and not error
}
landed: Final = eventually(
lambda: _landed(rig, settled), lambda seen: all(len(values) >= 1 for values in seen.values()), seconds=60
)
_assert_exactly_once(rig, settled, landed)
assert all(len(values) <= 1 for values in _landed(rig, served).values()), _landed(rig, served)
lost: Final = _landed(rig, {identity: sessions[index] for index, (identity, error) in enumerate(results) if error})
assert all(values == () for values in lost.values()), lost
@pytest.mark.covers("other.observability.otel.conversation_id_flushes_on_shutdown")
def test_terminating_the_proxy_right_after_a_burst_flushes_every_span_before_exit(
provider: Wire, collector: Collector, tmp_path_factory: pytest.TempPathFactory
) -> None:
factory: Final = RigFactory(provider, collector, tmp_path_factory.mktemp("otel-shutdown"), {}, 2)
started: Final = factory.start()
rig: Final = next(started)
sessions: Final = {index: f"conv-{index}-{uuid.uuid4().hex}" for index in range(10)}
markers: Final = tuple(_marker() for _ in range(10))
responses: Final = tuple(
rig.chat(markers[index], headers={"x-litellm-session-id": sessions[index]}) for index in range(10)
)
assert all(response.status_code == 200 for response in responses), [response.text for response in responses]
expected: Final = {f"chatcmpl-{markers[index]}": sessions[index] for index in range(10)}
drained: Final = eventually(
lambda: _landed(rig, expected), lambda seen: all(len(values) >= 1 for values in seen.values()), seconds=60
)
assert drained == {session: (identity,) for identity, session in expected.items()}, drained
rig.process.process.terminate()
assert rig.process.process.wait(timeout=40) in (0, -signal.SIGTERM)
assert _landed(rig, expected) == drained
with pytest.raises(httpx.ConnectError):
next(started)

View file

@ -23,6 +23,7 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( # noqa: E4
from opentelemetry.trace import SpanKind # noqa: E402
from opentelemetry.trace.status import StatusCode # noqa: E402
from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY # noqa: E402
from litellm.integrations.otel import ( # noqa: E402
GenAI,
LiteLLM,
@ -175,6 +176,64 @@ def test_async_log_success_event_emits_llm_call_span():
assert span.status.status_code is StatusCode.UNSET
def test_llm_call_span_carries_the_callers_conversation_id():
logger, exporter = _logger()
kwargs = {**_kwargs(), "litellm_params": {"litellm_session_id": "conv-42", "metadata": {}}}
_emit_llm(logger, kwargs)
(span,) = exporter.get_finished_spans()
assert span.attributes[GenAI.CONVERSATION_ID] == "conv-42"
def test_llm_call_span_without_a_caller_session_has_no_conversation_id():
"""The proxy stamps ``metadata.trace_id`` with the OTel trace id and
``get_litellm_params`` back-fills ``litellm_session_id`` from it."""
logger, exporter = _logger()
otel_trace_id = "6ca5745ef6780d958f62925747f7a5ee"
kwargs = {
**_kwargs(payload=_payload(trace_id=otel_trace_id)),
"litellm_trace_id": otel_trace_id,
"litellm_params": {
"litellm_session_id": otel_trace_id,
"litellm_trace_id": otel_trace_id,
"metadata": {"trace_id": otel_trace_id},
},
}
_emit_llm(logger, kwargs)
(span,) = exporter.get_finished_spans()
assert GenAI.CONVERSATION_ID not in span.attributes
def test_llm_call_span_keeps_the_header_session_when_the_proxy_generated_a_body_one():
"""``missing_session_id: generate`` mints a body session and marks it, but the
caller's ``langfuse_session_id`` header is still their conversation."""
logger, exporter = _logger()
kwargs = {
**_kwargs(),
"litellm_params": {
"litellm_session_id": "minted-by-proxy",
"metadata": {"session_id": "minted-by-proxy", SESSION_ID_GENERATED_METADATA_KEY: True},
"proxy_server_request": {"headers": {"langfuse_session_id": "conv-header"}},
},
}
_emit_llm(logger, kwargs)
(span,) = exporter.get_finished_spans()
assert span.attributes[GenAI.CONVERSATION_ID] == "conv-header"
def test_replayed_llm_call_span_does_not_take_the_payloads_session_id():
"""``/callback_logs`` replays a finished payload whose ``litellm_params`` hold
only key metadata; a session minted under ``missing_session_id: generate``
lands there without its marker, so ``payload.session_id`` is never trusted."""
logger, exporter = _logger()
kwargs = {
**_kwargs(payload=_payload(session_id="minted-then-replayed", trace_id="minted-then-replayed")),
"litellm_params": {"metadata": {"user_api_key_hash": "hsh"}},
}
_emit_llm(logger, kwargs)
(span,) = exporter.get_finished_spans()
assert GenAI.CONVERSATION_ID not in span.attributes
def test_streaming_span_carries_time_to_first_chunk():
logger, exporter = _logger()
kwargs = {

View file

@ -11,6 +11,7 @@ from typing import Final
import pytest
import litellm
from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY
from litellm.integrations.otel import (
BAGGAGE_PROMOTED_KEYS,
DB,
@ -1386,6 +1387,148 @@ def test_llm_span_data_carries_the_caller_trace_controls():
assert LLMCallSpanData.from_standard_logging_payload(_sample_payload()).trace == TraceControls()
@pytest.mark.parametrize(
("litellm_params", "expected"),
[
({"litellm_session_id": "conv-body"}, "conv-body"),
({"metadata": {"session_id": "conv-meta"}}, "conv-meta"),
({"litellm_metadata": {"session_id": "conv-anthropic"}}, "conv-anthropic"),
({"proxy_server_request": {"headers": {"langfuse_session_id": "conv-header"}}}, "conv-header"),
({"litellm_session_id": "conv-body", "metadata": {"session_id": "conv-meta"}}, "conv-body"),
({"litellm_session_id": "", "metadata": {"session_id": ""}}, None),
({"litellm_trace_id": "trace-only", "metadata": {"trace_id": "trace-only"}}, None),
(
{
"litellm_session_id": "0" * 32,
"litellm_trace_id": "0" * 32,
"metadata": {"trace_id": "0" * 32},
},
None,
),
(
{
"litellm_session_id": "0" * 32,
"litellm_trace_id": "0" * 32,
"metadata": {"trace_id": "0" * 32},
"proxy_server_request": {"headers": {"langfuse_session_id": "conv-header"}},
},
"conv-header",
),
(
{
"litellm_session_id": "minted-by-proxy",
"metadata": {"session_id": "minted-by-proxy", SESSION_ID_GENERATED_METADATA_KEY: True},
},
None,
),
(
{
"litellm_session_id": "minted-by-proxy",
"metadata": {"session_id": "minted-by-proxy", SESSION_ID_GENERATED_METADATA_KEY: True},
"proxy_server_request": {"headers": {"langfuse_session_id": "conv-header"}},
},
"conv-header",
),
(
{
"litellm_session_id": "minted-by-proxy",
"metadata": {"session_id": "conv-other-key"},
"litellm_metadata": {"session_id": "minted-by-proxy", SESSION_ID_GENERATED_METADATA_KEY: True},
},
"conv-other-key",
),
(
{
"litellm_session_id": "minted-by-proxy",
"metadata": {"session_id": "minted-by-proxy", SESSION_ID_GENERATED_METADATA_KEY: True},
"litellm_metadata": {"session_id": "conv-other-key"},
},
"conv-other-key",
),
(
{
"litellm_session_id": "conv-x-header",
"litellm_trace_id": "conv-x-header",
"metadata": {"trace_id": "conv-x-header", "session_id": "conv-x-header"},
},
"conv-x-header",
),
({}, None),
],
ids=[
"litellm_session_id",
"metadata",
"anthropic-metadata",
"langfuse-header",
"litellm_session_id-beats-metadata",
"blank-values",
"trace-id-is-not-a-session",
"backfilled-from-otel-trace-id-is-not-a-conversation",
"backfilled-trace-id-does-not-shadow-the-header",
"proxy-generated-is-not-a-conversation",
"proxy-generated-does-not-shadow-the-header",
"proxy-generated-on-litellm_metadata-does-not-shadow-metadata",
"proxy-generated-on-metadata-does-not-shadow-litellm_metadata",
"x-litellm-session-id-header-sets-trace-and-session",
"empty",
],
)
def test_llm_call_event_resolves_the_callers_conversation_id(litellm_params, expected):
kwargs: Final = {"litellm_params": litellm_params, "litellm_trace_id": "per-request-uuid"}
assert LLMCallEvent.from_dict(kwargs).session_id == expected
@pytest.mark.parametrize(
("litellm_params", "payload", "expected"),
[
(
{"metadata": {"user_api_key_hash": "hsh"}},
{"session_id": "minted-then-replayed", "trace_id": "minted-then-replayed"},
None,
),
(
{"metadata": {"user_api_key_hash": "hsh"}},
{"session_id": "conv-replayed", "trace_id": "0af7651916cd43dd8448eb211c80319c"},
None,
),
({"litellm_session_id": "conv-live"}, {"session_id": "conv-replayed"}, "conv-live"),
(
{
"litellm_session_id": "minted-by-proxy",
"metadata": {"session_id": "minted-by-proxy", SESSION_ID_GENERATED_METADATA_KEY: True},
},
{"session_id": "minted-by-proxy"},
None,
),
],
ids=[
"replayed-minted-session-stays-hidden",
"replayed-payload-is-not-a-source",
"live-params-win",
"generated-stays-hidden",
],
)
def test_llm_call_event_never_reads_the_replayed_payloads_session_id(litellm_params, payload, expected):
"""``/callback_logs`` rebuilds ``litellm_params`` with key metadata only, so a
``StandardLoggingPayload`` minted under ``missing_session_id: generate`` arrives
without its generated marker and is indistinguishable from a caller's session;
the payload is therefore never a source for the conversation id."""
kwargs: Final = {
"litellm_params": litellm_params,
"standard_logging_object": _sample_payload(**payload),
}
assert LLMCallEvent.from_dict(kwargs).session_id == expected
def test_llm_span_stamps_gen_ai_conversation_id_only_when_the_caller_sent_one():
with_session: Final = LLMCallSpanData.from_standard_logging_payload(_sample_payload(), session_id="conv-1")
assert GenAIMapper().map(with_session)[GenAI.CONVERSATION_ID] == "conv-1"
without: Final = LLMCallSpanData.from_standard_logging_payload(_sample_payload(trace_id="per-request-uuid"))
assert without.session_id is None
assert GenAI.CONVERSATION_ID not in GenAIMapper().map(without)
def test_llm_span_carries_proxy_request_route():
"""The LLM span records the proxy route the request arrived on, so it can be
filtered by endpoint (``/v1/responses`` vs ``/v1/chat/completions``) without

View file

@ -903,6 +903,61 @@ def test_authorize_post_accepts_ui_session_cookie(unauthenticated_client):
assert _byok_auth_codes[code]["user_id"] == "browser-user-42"
def test_authorize_post_rejects_cookie_with_revoked_session_key(unauthenticated_client):
"""The cookie JWT stays signature-valid until ``exp``, but logout /
password-change revocation deletes the DB-backed session key sealed
inside it. A cookie whose embedded key no longer resolves must not
authorize BYOK writes."""
import jwt as _jwt
with (
patch("litellm.proxy.proxy_server.master_key", "test-master-key"),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_key_object",
new=AsyncMock(side_effect=Exception("key not found")),
),
):
cookie_jwt = _jwt.encode(
{
"user_id": "browser-user-42",
"key": "sk-revoked-session-key",
"login_method": "sso",
"exp": int(time.time()) + 3600,
},
"test-master-key",
algorithm="HS256",
)
resp = _authorize_post_with_cookie(unauthenticated_client, cookie_jwt)
assert resp.status_code == 401
def test_authorize_post_accepts_cookie_with_live_session_key(unauthenticated_client):
"""A cookie whose embedded session key still resolves keeps working."""
import jwt as _jwt
with (
patch("litellm.proxy.proxy_server.master_key", "test-master-key"),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_key_object",
new=AsyncMock(return_value=UserAPIKeyAuth(user_id="browser-user-42")),
),
):
cookie_jwt = _jwt.encode(
{
"user_id": "browser-user-42",
"key": "sk-live-session-key",
"login_method": "sso",
"exp": int(time.time()) + 3600,
},
"test-master-key",
algorithm="HS256",
)
resp = _authorize_post_with_cookie(unauthenticated_client, cookie_jwt)
assert resp.status_code == 302
def test_authorize_post_rejects_cookie_signed_with_wrong_key(unauthenticated_client):
"""A cookie JWT signed with a different key than the proxy's master_key
must not grant access — otherwise an attacker who can forge a JWT

View file

@ -2137,7 +2137,7 @@ class TestPasswordResetRequiredSessionMinting:
row = _db_user_row(password="Str0ng!Passw0rd", password_reset_required=True)
result, key_kwargs = await self._login(_prisma_with_user(row))
assert key_kwargs["allowed_routes"] == ["/user/password/change"]
assert key_kwargs["allowed_routes"] == ["/user/password/change", "/session/logout"]
assert key_kwargs["metadata"] == {"login_method": "username_password", "password_reset_required": True}
assert result.password_reset_required is True
@ -2198,7 +2198,7 @@ class TestPasswordResetRequiredSessionMinting:
result, key_kwargs, _ = await self._login_with_screen_result(mock_prisma_client, breached=True)
assert key_kwargs["allowed_routes"] == ["/user/password/change"]
assert key_kwargs["allowed_routes"] == ["/user/password/change", "/session/logout"]
assert key_kwargs["metadata"] == {"login_method": "username_password", "password_reset_required": True}
assert result.password_reset_required is True

View file

@ -477,6 +477,72 @@ async def test_claim_token_sets_accepted_at_after_password_written():
assert outer_claims["key"] == "sk-generated-key"
@pytest.mark.asyncio
async def test_claim_token_revokes_existing_ui_sessions():
"""A claimed invite/reset link changes the password; any UI session minted
under the old password may be in hostile hands and must be revoked. The
sweep runs before the fresh session key is minted, so revoke-all is safe."""
from litellm.proxy.proxy_server import claim_onboarding_link
invite = _make_invite(is_accepted=False)
user = _make_user()
prisma = _make_prisma(invite, user)
request = _make_claim_request(_make_onboarding_token())
data = InvitationClaim(
invitation_link="invite-abc",
user_id="user-123",
password="NewP@ssw0rd123",
)
mock_token_response = {"token": "sk-generated-key", "user_id": "user-123"}
revoke_mock = AsyncMock(return_value=1)
mint_order: list[str] = []
async def _mint(*args, **kwargs):
mint_order.append("mint")
return mock_token_response
async def _revoke(*args, **kwargs):
mint_order.append("revoke")
return 1
revoke_mock.side_effect = _revoke
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.master_key", "sk-test"),
patch( # test-quality-ok: claim_onboarding_link reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.general_settings", _POLICY_NO_BREACH_CHECK
),
patch("litellm.proxy.proxy_server.premium_user", False),
patch(
"litellm.proxy.proxy_server.generate_key_helper_fn",
new_callable=AsyncMock,
side_effect=_mint,
),
patch(
"litellm.proxy.management_endpoints.session_endpoints.revoke_ui_session_keys",
revoke_mock,
),
patch(
"litellm.proxy.proxy_server.get_custom_url",
return_value="http://localhost:4000/",
),
patch(
"litellm.proxy.proxy_server.get_disabled_non_admin_personal_key_creation",
return_value=False,
),
patch("litellm.proxy.proxy_server.get_server_root_path", return_value=""),
):
await claim_onboarding_link(data=data, request=request)
revoke_mock.assert_awaited_once()
assert revoke_mock.await_args.kwargs["user_id"] == "user-123"
# The sweep must precede the mint or it would kill the fresh session too.
assert mint_order == ["revoke", "mint"]
@pytest.mark.asyncio
async def test_claim_token_rolls_back_invite_when_session_key_mint_fails():
"""A session key failure must not leave the invite permanently consumed."""

View file

@ -4905,3 +4905,69 @@ async def test_delete_user_writes_deleted_audit_log_for_user_keys(mocker):
assert audit_row.object_id == user_key.token
assert audit_row.changed_by
assert json.loads(audit_row.before_value)["token"] == user_key.token
@pytest.mark.asyncio
async def test_user_update_password_revokes_target_sessions(_admin_prisma, mocker):
"""An admin-set password implies the old one may be compromised: every UI
session belonging to the target user must be revoked after the write."""
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_update_single_user_helper,
)
mocker.patch( # test-quality-ok: same module-global mocking every test in this file already uses
"litellm.proxy.proxy_server.general_settings",
{"password_policy_check_breached_passwords": False},
)
mock_prisma_client = _admin_prisma
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {"user_id": "target-user"}
existing_user.user_id = "target-user"
mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=existing_user)
mock_prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": "target-user"})
mock_prisma_client.jsonify_object = mocker.MagicMock(side_effect=lambda x: x)
revoke_mock = mocker.patch(
"litellm.proxy.management_endpoints.session_endpoints.revoke_ui_session_keys",
new=mocker.AsyncMock(return_value=2),
)
user_request = UpdateUserRequest(user_id="target-user", password="Str0ng!Passw0rd")
admin_caller = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN)
await _update_single_user_helper(user_request=user_request, user_api_key_dict=admin_caller)
revoke_mock.assert_awaited_once()
revoke_kwargs = revoke_mock.await_args.kwargs
assert revoke_kwargs["user_id"] == "target-user"
# Revoke-all: the admin's own session is not among the target's sessions.
assert revoke_kwargs.get("keep_hashed_token") is None
@pytest.mark.asyncio
async def test_user_update_without_password_revokes_nothing(_admin_prisma, mocker):
"""A non-password /user/update must not touch the target's sessions."""
from litellm.proxy.management_endpoints.internal_user_endpoints import (
_update_single_user_helper,
)
mock_prisma_client = _admin_prisma
existing_user = mocker.MagicMock()
existing_user.model_dump.return_value = {"user_id": "target-user"}
existing_user.user_id = "target-user"
mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=existing_user)
mock_prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": "target-user"})
mock_prisma_client.jsonify_object = mocker.MagicMock(side_effect=lambda x: x)
revoke_mock = mocker.patch(
"litellm.proxy.management_endpoints.session_endpoints.revoke_ui_session_keys",
new=mocker.AsyncMock(return_value=0),
)
user_request = UpdateUserRequest(user_id="target-user", user_email="new@example.com")
admin_caller = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN)
await _update_single_user_helper(user_request=user_request, user_api_key_dict=admin_caller)
revoke_mock.assert_not_awaited()

View file

@ -380,6 +380,73 @@ async def test_change_password_failure_emits_no_audit_log():
audit_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_change_password_revokes_other_sessions_keeping_callers():
"""A successful change revokes the user's other UI sessions (the old
password may be compromised) while keeping the session that just proved
it holds the current password."""
from litellm.proxy._types import ChangePasswordRequest
prisma = _make_prisma(_make_user_row(hash_password(CURRENT_PASSWORD)))
revoke_mock = AsyncMock(return_value=0)
caller = UserAPIKeyAuth(
user_id="user-123",
token="hashed-caller-token",
team_id=UI_TEAM_ID,
metadata=dict(PASSWORD_SESSION_METADATA),
)
with (
patch( # test-quality-ok: change_password reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.prisma_client", prisma
),
patch( # test-quality-ok: change_password reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.general_settings", _POLICY_NO_BREACH_CHECK
),
patch(
"litellm.proxy.management_endpoints.password_endpoints.revoke_ui_session_keys",
revoke_mock,
),
):
await change_password(
data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD),
user_api_key_dict=caller,
)
revoke_mock.assert_awaited_once()
revoke_kwargs = revoke_mock.await_args.kwargs
assert revoke_kwargs["user_id"] == "user-123"
assert revoke_kwargs["keep_hashed_token"] == "hashed-caller-token"
@pytest.mark.asyncio
async def test_change_password_failure_revokes_no_sessions():
from litellm.proxy._types import ChangePasswordRequest
prisma = _make_prisma(_make_user_row(hash_password(CURRENT_PASSWORD)))
revoke_mock = AsyncMock(return_value=0)
with (
patch( # test-quality-ok: change_password reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.prisma_client", prisma
),
patch( # test-quality-ok: change_password reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.general_settings", _POLICY_NO_BREACH_CHECK
),
patch(
"litellm.proxy.management_endpoints.password_endpoints.revoke_ui_session_keys",
revoke_mock,
),
):
with pytest.raises(HTTPException):
await change_password(
data=ChangePasswordRequest(current_password="not-the-password", new_password=NEW_PASSWORD),
user_api_key_dict=_caller(),
)
revoke_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_change_password_requires_db():
from litellm.proxy._types import ChangePasswordRequest

View file

@ -0,0 +1,300 @@
"""
Tests for POST /session/logout and revoke_ui_session_keys
(litellm/proxy/management_endpoints/session_endpoints.py).
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException, Response
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
from litellm.proxy._types import LiteLLM_VerificationToken, UserAPIKeyAuth
from litellm.proxy.management_endpoints.session_endpoints import (
revoke_ui_session_keys,
session_logout,
)
HASHED_TOKEN = "hashed-session-token"
USER_ID = "user-123"
def _session_row(token: str = HASHED_TOKEN, user_id: str = USER_ID) -> LiteLLM_VerificationToken:
return LiteLLM_VerificationToken(token=token, team_id=UI_SESSION_TOKEN_TEAM_ID, user_id=user_id)
def _make_prisma(
find_unique_row: LiteLLM_VerificationToken | None = None,
find_many_rows: list[LiteLLM_VerificationToken] | None = None,
) -> MagicMock:
prisma = MagicMock()
table = prisma.db.litellm_verificationtoken
table.find_unique = AsyncMock(return_value=find_unique_row)
table.find_many = AsyncMock(return_value=find_many_rows or [])
table.delete_many = AsyncMock(return_value=1)
return prisma
def _ui_session_caller(token: str | None = HASHED_TOKEN) -> UserAPIKeyAuth:
return UserAPIKeyAuth(token=token, team_id=UI_SESSION_TOKEN_TEAM_ID, user_id=USER_ID)
def _patched_globals(prisma):
return (
patch( # test-quality-ok: endpoint reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.prisma_client", prisma
),
patch( # test-quality-ok: endpoint reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.proxy_logging_obj", None
),
patch( # test-quality-ok: endpoint reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.user_api_key_cache", MagicMock()
),
)
@pytest.mark.asyncio
async def test_session_logout_revokes_presented_session():
prisma = _make_prisma(find_unique_row=_session_row())
persist_mock = AsyncMock()
evict_mock = AsyncMock()
p1, p2, p3 = _patched_globals(prisma)
with (
p1,
p2,
p3,
patch(
"litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens",
persist_mock,
),
patch(
"litellm.proxy.management_endpoints.session_endpoints.delete_cache_key_objects",
evict_mock,
),
):
response = await session_logout(
response=Response(),
user_api_key_dict=_ui_session_caller(),
)
assert response.message == "Session revoked."
delete_kwargs = prisma.db.litellm_verificationtoken.delete_many.call_args.kwargs
assert delete_kwargs["where"] == {"token": HASHED_TOKEN}
# Audit record persisted before the row is gone.
persist_mock.assert_awaited_once()
assert persist_mock.await_args.kwargs["keys"][0].token == HASHED_TOKEN
# Cache evicted + broadcast even on the delete path.
evict_mock.assert_awaited_once()
assert tuple(evict_mock.await_args.kwargs["hashed_tokens"]) == (HASHED_TOKEN,)
@pytest.mark.asyncio
async def test_session_logout_clears_token_cookie():
prisma = _make_prisma(find_unique_row=_session_row())
fastapi_response = Response()
p1, p2, p3 = _patched_globals(prisma)
with (
p1,
p2,
p3,
patch(
"litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens",
AsyncMock(),
),
patch(
"litellm.proxy.management_endpoints.session_endpoints.delete_cache_key_objects",
AsyncMock(),
),
):
await session_logout(
response=fastapi_response,
user_api_key_dict=_ui_session_caller(),
)
set_cookie_headers = [v.decode() for k, v in fastapi_response.raw_headers if k == b"set-cookie"]
assert any(h.startswith('token="";') or h.startswith("token=;") for h in set_cookie_headers)
@pytest.mark.asyncio
async def test_session_logout_refuses_non_ui_session_key():
"""The endpoint must not become a generic key-deletion oracle: a normal
virtual key (no UI team id) is refused outright."""
prisma = _make_prisma()
p1, p2, p3 = _patched_globals(prisma)
with p1, p2, p3:
with pytest.raises(HTTPException) as exc_info:
await session_logout(
response=Response(),
user_api_key_dict=UserAPIKeyAuth(token=HASHED_TOKEN, team_id="some-real-team", user_id=USER_ID),
)
assert exc_info.value.status_code == 403
prisma.db.litellm_verificationtoken.delete_many.assert_not_called()
@pytest.mark.asyncio
async def test_session_logout_is_idempotent_when_row_already_gone():
prisma = _make_prisma(find_unique_row=None)
evict_mock = AsyncMock()
p1, p2, p3 = _patched_globals(prisma)
with (
p1,
p2,
p3,
patch(
"litellm.proxy.management_endpoints.session_endpoints.delete_cache_key_objects",
evict_mock,
),
):
response = await session_logout(
response=Response(),
user_api_key_dict=_ui_session_caller(),
)
assert response.message == "Session already revoked."
prisma.db.litellm_verificationtoken.delete_many.assert_not_called()
# The cache entry may outlive the row; evict regardless.
evict_mock.assert_awaited_once()
@pytest.mark.asyncio
async def test_session_logout_requires_db():
p2 = patch("litellm.proxy.proxy_server.proxy_logging_obj", None)
p3 = patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
with (
patch( # test-quality-ok: endpoint reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.prisma_client", None
),
p2,
p3,
):
with pytest.raises(HTTPException) as exc_info:
await session_logout(
response=Response(),
user_api_key_dict=_ui_session_caller(),
)
assert exc_info.value.status_code == 500
@pytest.mark.asyncio
async def test_revoke_ui_session_keys_revokes_all_and_broadcasts():
rows = [_session_row(token="t1"), _session_row(token="t2"), _session_row(token="t3")]
prisma = _make_prisma(find_many_rows=rows)
persist_mock = AsyncMock()
evict_mock = AsyncMock()
p1, p2, p3 = _patched_globals(prisma)
with (
p1,
p2,
p3,
patch(
"litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens",
persist_mock,
),
patch(
"litellm.proxy.management_endpoints.session_endpoints.delete_cache_key_objects",
evict_mock,
),
):
revoked = await revoke_ui_session_keys(
user_id=USER_ID,
user_api_key_dict=_ui_session_caller(),
)
assert revoked == 3
find_kwargs = prisma.db.litellm_verificationtoken.find_many.call_args.kwargs
assert find_kwargs["where"] == {"user_id": USER_ID, "team_id": UI_SESSION_TOKEN_TEAM_ID}
delete_kwargs = prisma.db.litellm_verificationtoken.delete_many.call_args.kwargs
assert delete_kwargs["where"] == {"token": {"in": ["t1", "t2", "t3"]}}
persist_mock.assert_awaited_once()
evict_mock.assert_awaited_once()
assert evict_mock.await_args.kwargs["hashed_tokens"] == ["t1", "t2", "t3"]
@pytest.mark.asyncio
async def test_revoke_ui_session_keys_keeps_callers_session():
rows = [_session_row(token="t1"), _session_row(token=HASHED_TOKEN), _session_row(token="t3")]
prisma = _make_prisma(find_many_rows=rows)
p1, p2, p3 = _patched_globals(prisma)
with (
p1,
p2,
p3,
patch(
"litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens",
AsyncMock(),
),
patch(
"litellm.proxy.management_endpoints.session_endpoints.delete_cache_key_objects",
AsyncMock(),
),
):
revoked = await revoke_ui_session_keys(
user_id=USER_ID,
user_api_key_dict=_ui_session_caller(),
keep_hashed_token=HASHED_TOKEN,
)
assert revoked == 2
delete_kwargs = prisma.db.litellm_verificationtoken.delete_many.call_args.kwargs
assert delete_kwargs["where"] == {"token": {"in": ["t1", "t3"]}}
@pytest.mark.asyncio
async def test_revoke_ui_session_keys_noop_when_no_sessions():
prisma = _make_prisma(find_many_rows=[])
p1, p2, p3 = _patched_globals(prisma)
with p1, p2, p3:
revoked = await revoke_ui_session_keys(
user_id=USER_ID,
user_api_key_dict=_ui_session_caller(),
)
assert revoked == 0
prisma.db.litellm_verificationtoken.delete_many.assert_not_called()
@pytest.mark.asyncio
async def test_revoke_ui_session_keys_failure_is_swallowed():
"""The password write has already committed when this runs; a revocation
failure must not fail the caller's request."""
prisma = _make_prisma(find_many_rows=[_session_row(token="t1")])
prisma.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=RuntimeError("db down"))
p1, p2, p3 = _patched_globals(prisma)
with (
p1,
p2,
p3,
patch(
"litellm.proxy.management_endpoints.session_endpoints._persist_deleted_verification_tokens",
AsyncMock(),
),
):
revoked = await revoke_ui_session_keys(
user_id=USER_ID,
user_api_key_dict=_ui_session_caller(),
)
assert revoked == 0
@pytest.mark.asyncio
async def test_revoke_ui_session_keys_noop_without_db():
with patch( # test-quality-ok: helper reads proxy_server module globals; no injection seam
"litellm.proxy.proxy_server.prisma_client", None
):
revoked = await revoke_ui_session_keys(
user_id=USER_ID,
user_api_key_dict=_ui_session_caller(),
)
assert revoked == 0

View file

@ -15,7 +15,7 @@ import { changePasswordCall, getProxyBaseUrl } from "@/components/networking";
import { extractProxyErrorMessage } from "@/lib/http/client";
import { useZodForm } from "@/lib/forms/useZodForm";
import { toast } from "@/lib/toast";
import { clearTokenCookies } from "@/utils/cookieUtils";
import { revokeSessionAndClearClientState } from "@/app/(dashboard)/hooks/useLogout";
import { getLoginUrl } from "@/utils/returnUrlUtils";
const changePasswordSchema = z
@ -47,8 +47,10 @@ export function ChangePasswordForm() {
await changePasswordCall(accessToken, values.currentPassword, values.newPassword);
if (passwordResetRequired) {
// The session key was minted restricted; only a fresh login lifts it.
// Revoke it server-side too (best-effort) so it doesn't sit valid
// until the expiry reaper gets to it.
toast.success("Password updated. Please log in with your new password.");
clearTokenCookies();
await revokeSessionAndClearClientState(accessToken);
window.location.replace(getLoginUrl(getProxyBaseUrl()));
return;
}

View file

@ -0,0 +1,72 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
const sessionLogoutCall = vi.hoisted(() => vi.fn());
const clearTokenCookies = vi.hoisted(() => vi.fn());
const clearStoredReturnUrl = vi.hoisted(() => vi.fn());
vi.mock("@/components/networking", () => ({
sessionLogoutCall,
}));
vi.mock("@/utils/cookieUtils", () => ({
clearTokenCookies,
}));
vi.mock("@/utils/returnUrlUtils", () => ({
clearStoredReturnUrl,
}));
vi.mock("@/app/(dashboard)/hooks/proxySettings/useProxySettings", () => ({
default: vi.fn(() => ({ PROXY_LOGOUT_URL: "" })),
}));
import { revokeSessionAndClearClientState } from "./useLogout";
describe("revokeSessionAndClearClientState", () => {
beforeEach(() => {
vi.clearAllMocks();
sessionLogoutCall.mockResolvedValue({ message: "Session revoked." });
localStorage.setItem("litellm_selected_worker_id", "w1");
localStorage.setItem("litellm_worker_url", "https://worker.example");
});
it("revokes the session server-side before clearing the token cookie", async () => {
const order: string[] = [];
sessionLogoutCall.mockImplementation(async () => {
order.push("revoke");
return { message: "Session revoked." };
});
clearTokenCookies.mockImplementation(() => {
order.push("clearCookies");
});
await revokeSessionAndClearClientState("sk-token");
expect(sessionLogoutCall).toHaveBeenCalledWith("sk-token");
// The cookie holds the credential that authenticates the revoke call, so
// clearing it first would orphan the server-side key.
expect(order).toEqual(["revoke", "clearCookies"]);
});
it("clears all client state", async () => {
await revokeSessionAndClearClientState("sk-token");
expect(clearTokenCookies).toHaveBeenCalled();
expect(clearStoredReturnUrl).toHaveBeenCalled();
expect(localStorage.getItem("litellm_selected_worker_id")).toBeNull();
expect(localStorage.getItem("litellm_worker_url")).toBeNull();
});
it("still clears client state when the revoke call rejects", async () => {
sessionLogoutCall.mockRejectedValue(new Error("proxy unreachable"));
await revokeSessionAndClearClientState("sk-token");
expect(clearTokenCookies).toHaveBeenCalled();
expect(localStorage.getItem("litellm_selected_worker_id")).toBeNull();
});
it("skips the server call without a token but still clears client state", async () => {
await revokeSessionAndClearClientState(null);
expect(sessionLogoutCall).not.toHaveBeenCalled();
expect(clearTokenCookies).toHaveBeenCalled();
});
});

View file

@ -1,7 +1,29 @@
import { sessionLogoutCall } from "@/components/networking";
import { clearTokenCookies } from "@/utils/cookieUtils";
import { clearStoredReturnUrl } from "@/utils/returnUrlUtils";
import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings";
/**
* Revokes the session key server-side, then clears client state. Exported for
* flows that navigate somewhere other than PROXY_LOGOUT_URL (worker switch,
* forced password reset). The server call must happen BEFORE the cookies are
* cleared (the token authenticates it) and is best-effort: local logout must
* still complete when the server is unreachable.
*/
export async function revokeSessionAndClearClientState(accessToken: string | null): Promise<void> {
if (accessToken) {
try {
await sessionLogoutCall(accessToken);
} catch {
// Best-effort: the key still expires server-side at its session TTL.
}
}
clearTokenCookies();
clearStoredReturnUrl();
localStorage.removeItem("litellm_selected_worker_id");
localStorage.removeItem("litellm_worker_url");
}
/**
* Shared sign-out handler. Used by both the top navbar and the sidebar footer so
* the two entry points can never drift on which client state gets cleared.
@ -10,10 +32,8 @@ export function useLogout(accessToken: string | null): () => void {
const proxySettings = useProxySettings(accessToken);
return () => {
clearTokenCookies();
clearStoredReturnUrl();
localStorage.removeItem("litellm_selected_worker_id");
localStorage.removeItem("litellm_worker_url");
window.location.href = proxySettings.PROXY_LOGOUT_URL || "";
void revokeSessionAndClearClientState(accessToken).finally(() => {
window.location.href = proxySettings.PROXY_LOGOUT_URL || "";
});
};
}

View file

@ -5,9 +5,8 @@ import { useWorker } from "@/hooks/useWorker";
import { getProxyBaseUrl } from "@/components/networking";
import { uiHref } from "@/utils/uiHref";
import { useTheme } from "@/contexts/ThemeContext";
import { clearTokenCookies } from "@/utils/cookieUtils";
import { clearStoredReturnUrl, getLoginUrl } from "@/utils/returnUrlUtils";
import useProxySettings from "@/app/(dashboard)/hooks/proxySettings/useProxySettings";
import { revokeSessionAndClearClientState, useLogout } from "@/app/(dashboard)/hooks/useLogout";
import { getLoginUrl } from "@/utils/returnUrlUtils";
import { Badge } from "@/components/ui/badge";
import { PanelLeftClose, PanelLeftOpen } from "lucide-react";
import Link from "next/link";
@ -38,7 +37,6 @@ const Navbar: React.FC<NavbarProps> = ({
onToggleSidebar,
}) => {
const baseUrl = getProxyBaseUrl();
const proxySettings = useProxySettings(accessToken);
const { logoUrl } = useTheme();
const { data: healthData } = useHealthReadinessDetails(accessToken);
const version = healthData?.litellm_version;
@ -50,19 +48,12 @@ const Navbar: React.FC<NavbarProps> = ({
const imageUrl = logoUrl || `${baseUrl}/get_image`;
const darkImageUrl = logoUrl || `${baseUrl}/get_image?theme=dark`;
const handleLogout = () => {
clearTokenCookies();
localStorage.removeItem("litellm_selected_worker_id");
localStorage.removeItem("litellm_worker_url");
window.location.href = proxySettings.PROXY_LOGOUT_URL || "";
};
const handleLogout = useLogout(accessToken);
const handleWorkerSwitch = (workerId: string) => {
clearTokenCookies();
clearStoredReturnUrl();
localStorage.removeItem("litellm_selected_worker_id");
localStorage.removeItem("litellm_worker_url");
window.location.href = `${getLoginUrl()}?worker=${encodeURIComponent(workerId)}`;
void revokeSessionAndClearClientState(accessToken).finally(() => {
window.location.href = `${getLoginUrl()}?worker=${encodeURIComponent(workerId)}`;
});
};
return (

View file

@ -1595,6 +1595,18 @@ export const claimOnboardingToken = async (
}
};
/**
* Revokes the UI session key server-side (POST /session/logout). Best-effort
* with a short timeout: logout must still complete locally when the server is
* unreachable, so callers swallow rejections.
*/
export const sessionLogoutCall = async (accessToken: string): Promise<{ message: string }> => {
return await apiClient.post(`/session/logout`, {
accessToken,
signal: AbortSignal.timeout(3000),
});
};
export const changePasswordCall = async (
accessToken: string,
currentPassword: string,

View file

@ -14485,6 +14485,31 @@ export interface paths {
patch?: never;
trace?: never;
};
"/session/logout": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
put?: never;
/**
* Session Logout
* @description Revoke the UI session key this request authenticated with.
*
* Only accepts UI session keys (minted by dashboard login); any other
* credential is refused, so this can never be used to delete arbitrary keys.
* Revokes only the presented session, not the user's other sessions.
* Idempotent: logging out an already-revoked session succeeds.
*/
post: operations["session_logout_session_logout_post"];
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/settings": {
parameters: {
query?: never;
@ -38133,6 +38158,11 @@ export interface components {
/** Timeout */
timeout?: number | null;
};
/** SessionLogoutResponse */
SessionLogoutResponse: {
/** Message */
message: string;
};
/**
* ShadowEvalJobResponse
* @description A shadow-eval job over one or more targets, each with its own budget and stop state;
@ -60958,6 +60988,26 @@ export interface operations {
};
};
};
session_logout_session_logout_post: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody?: never;
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["SessionLogoutResponse"];
};
};
};
};
active_callbacks_settings_get: {
parameters: {
query?: never;