mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
chore: merge origin/main into litellm_langfuse_otel_end_user_id
This commit is contained in:
commit
e4b5f26a77
35 changed files with 2169 additions and 234 deletions
186
.github/workflows/create-release.yml
vendored
186
.github/workflows/create-release.yml
vendored
|
|
@ -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 }}
|
||||
|
|
@ -26,6 +26,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/v2/login",
|
||||
"/v3/login",
|
||||
"/logout",
|
||||
"/session/logout",
|
||||
"/token",
|
||||
"/onboarding/",
|
||||
"/audit",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
175
litellm/proxy/management_endpoints/session_endpoints.py
Normal file
175
litellm/proxy/management_endpoints/session_endpoints.py
Normal 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.",
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
830
tests/integration/observability/test_otel_conversation_id.py
Normal file
830
tests/integration/observability/test_otel_conversation_id.py
Normal 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)
|
||||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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 || "";
|
||||
});
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
50
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
50
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue