Merge remote-tracking branch 'origin/main' into litellm_team_member_budget_alerts_merge

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

# Conflicts:
#	tests/integration/contracts.json
This commit is contained in:
ryan 2026-09-23 08:10:27 +00:00
commit 35908e015b
102 changed files with 9545 additions and 481 deletions

View file

@ -9,7 +9,6 @@ fi
suite="${1:?integration suite required}"
results="test-results/integration-${suite}"
mkdir -p "$results"
shard_timeout=11m
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
upstream_pid=""
proxy_pid=""
@ -131,6 +130,7 @@ start_proxy() {
fi
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
@ -176,7 +176,7 @@ if [ "$suite" = browser ]; then
exit 0
fi
timeout --signal=TERM --kill-after=20s "$shard_timeout" env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
INTEGRATION_PROXY_URL="$INTEGRATION_PROXY_URL" INTEGRATION_PEER_URL="$INTEGRATION_PEER_URL" \

View file

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

View file

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

View file

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

View file

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

View file

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

File diff suppressed because it is too large Load diff

View file

@ -61,3 +61,4 @@ class LiteLLM_ManagedVectorStoresTable(LiteLLMPydanticObjectBase):
litellm_params: dict[str, Any] | None = None
team_id: str | None = None
user_id: str | None = None
is_config: bool = False

View file

@ -45790,6 +45790,10 @@
"title": "Custom Llm Provider",
"type": "string"
},
"is_config": {
"title": "Is Config",
"type": "boolean"
},
"litellm_credential_name": {
"anyOf": [
{
@ -45962,6 +45966,11 @@
"title": "Custom Llm Provider",
"type": "string"
},
"is_config": {
"default": false,
"title": "Is Config",
"type": "boolean"
},
"litellm_credential_name": {
"anyOf": [
{
@ -46203,7 +46212,7 @@
"paths": {
"/v1/vector_store/list": {
"get": {
"description": "List all available vector stores with optional filtering and pagination.\nCombines both in-memory vector stores and those stored in the database.\nDatabase is the source of truth - deleted stores are removed from memory, updated stores sync to memory.\n\nParameters:\n- page: int - Page number for pagination (default: 1)\n- page_size: int - Number of items per page (default: 100)",
"description": "List all available vector stores with optional filtering and pagination.\nCombines both in-memory vector stores and those stored in the database.\nDatabase is the source of truth for stores it owns: deleted stores are removed from memory, updated stores\nsync to memory. Stores declared in the config file are owned by the config file, are always listed, and are\nnever overwritten by database rows.\n\nParameters:\n- page: int - Page number for pagination (default: 1)\n- page_size: int - Number of items per page (default: 100)",
"operationId": "list_vector_stores_v1_vector_store_list_get",
"parameters": [
{
@ -46354,7 +46363,7 @@
},
"/vector_store/list": {
"get": {
"description": "List all available vector stores with optional filtering and pagination.\nCombines both in-memory vector stores and those stored in the database.\nDatabase is the source of truth - deleted stores are removed from memory, updated stores sync to memory.\n\nParameters:\n- page: int - Page number for pagination (default: 1)\n- page_size: int - Number of items per page (default: 100)",
"description": "List all available vector stores with optional filtering and pagination.\nCombines both in-memory vector stores and those stored in the database.\nDatabase is the source of truth for stores it owns: deleted stores are removed from memory, updated stores\nsync to memory. Stores declared in the config file are owned by the config file, are always listed, and are\nnever overwritten by database rows.\n\nParameters:\n- page: int - Page number for pagination (default: 1)\n- page_size: int - Number of items per page (default: 100)",
"operationId": "list_vector_stores_vector_store_list_get",
"parameters": [
{

View file

@ -18,6 +18,10 @@ if TYPE_CHECKING:
AUTH_CACHE_INVALIDATION_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation"
_POLL_TIMEOUT_SECONDS: Final = 1.0
_MAX_PENDING_PUBLISHES: Final = 1024
_MAX_IN_FLIGHT_PUBLISHES: Final = 16
_pending_publishes: Final[set[asyncio.Task[None]]] = set() # mutable-ok: strong refs keep background publishes alive
_in_flight_publishes: Final = asyncio.Semaphore(_MAX_IN_FLIGHT_PUBLISHES)
_BACKOFF_INITIAL_SECONDS: Final = 5.0
_BACKOFF_MAX_SECONDS: Final = 60.0
@ -67,6 +71,21 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None:
)
async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None:
try:
client: Final = _pubsub_capable_client(redis_cache)
if client is None:
verbose_proxy_logger.debug(
"auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support",
cache_key,
)
return
async with _in_flight_publishes:
await client.publish(auth_cache_invalidation_channel(redis_cache), message)
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e)
async def publish_auth_cache_invalidation(
cache_key: str, new_value: float | None = None, ttl: float | None = None
) -> None:
@ -80,24 +99,34 @@ async def publish_auth_cache_invalidation(
writes the value into its additional in-memory caches rather than deleting
the key. A spend reset uses this so the handler's self-delivered message
cannot erase the freshly-written post-reset counter or floor marker.
The Redis round trip runs as a background task: this call returns once the
publish has been handed to the event loop, so a Redis that accepts
connections but never replies costs the caller nothing. The DB write has
already committed and the local eviction already happened, so the caller
has nothing to do with the publish result. At most 16 publishes hold a
Redis connection at once; the rest wait in the task set, so a wedge cannot
drain the shared connection pool.
"""
redis_cache: Final = coordination_redis_cache()
if redis_cache is None:
return
try:
client: Final = _pubsub_capable_client(redis_cache)
if client is None:
verbose_proxy_logger.debug(
"auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support",
cache_key,
)
return
await client.publish(
auth_cache_invalidation_channel(redis_cache),
_cache_invalidation_message_json(cache_key, new_value=new_value, ttl=ttl),
_pending_publishes.difference_update({task for task in _pending_publishes if task.done()})
if len(_pending_publishes) >= _MAX_PENDING_PUBLISHES:
verbose_proxy_logger.warning(
"auth cache invalidation publish for %s dropped: %d publishes already waiting on redis; "
"other workers keep their cached copy until its TTL expires",
cache_key,
len(_pending_publishes),
)
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
verbose_proxy_logger.warning("auth cache invalidation publish for %s failed: %s", cache_key, e)
return
task: Final = asyncio.create_task(
_publish_to_redis(
redis_cache, cache_key, _cache_invalidation_message_json(cache_key, new_value=new_value, ttl=ttl)
)
)
_pending_publishes.add(task)
await asyncio.sleep(0)
async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "UserApiKeyCache") -> None:
@ -106,8 +135,8 @@ async def evict_and_broadcast(cache_keys: Sequence[str], user_api_key_cache: "Us
Every endpoint that mutates a cached object must call this: auth serves those objects
cache-first with no freshness check, so a mutation that leaves the entry in place keeps the
stale object enforced until its TTL expires (LIT-3803). Best-effort on both steps: the DB write
has already committed, so a cache backend error must not fail the endpoint.
stale object enforced until its TTL expires (LIT-3803). Best-effort: the DB write has already
committed, so a cache backend error must not fail the endpoint.
"""
for cache_key in cache_keys:
try:

View file

@ -12,6 +12,7 @@ from typing import Final
from litellm._logging import verbose_proxy_logger
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext, PolicyScope
@ -136,14 +137,46 @@ class PolicyMatcher:
context: PolicyMatchContext,
policies: dict[str, Policy] | None = None,
) -> Callable[[str], bool]:
"""Predicate telling whether a policy exists and its condition matches the context."""
"""
Predicate telling whether a policy exists and any policy in its
inheritance chain applies to the context. Admissions where the
policy's own condition missed but an ancestor applies are logged at
INFO, once per attachment scan.
"""
resolved: Final = policies if policies is not None else PolicyMatcher._registry_policies()
return lambda policy_name: bool(
PolicyMatcher.get_policies_with_matching_conditions(
policy_names=(policy_name,),
context=context,
policies=resolved,
def applies(policy_name: str) -> bool:
applying: Final = PolicyMatcher._applying_chain_members(
policy_name=policy_name, context=context, policies=resolved
)
if applying and policy_name not in applying:
verbose_proxy_logger.info(
"Policy '%s' applied through ancestor '%s' although its own condition did not match "
"(team_alias=%s, key_alias=%s, model=%s)",
policy_name,
applying[0],
context.team_alias,
context.key_alias,
context.model,
)
return bool(applying)
return applies
@staticmethod
def _applying_chain_members(
policy_name: str,
context: PolicyMatchContext,
policies: dict[str, Policy],
) -> tuple[str, ...]:
from litellm.proxy.policy_engine.condition_evaluator import ConditionEvaluator
chain: Final = PolicyResolver.resolve_inheritance_chain(policy_name=policy_name, policies=policies)
return tuple(
name
for name in chain
if (policy := policies.get(name)) is not None
and (policy.condition is None or ConditionEvaluator.evaluate(policy.condition, context))
)
@staticmethod
@ -160,11 +193,14 @@ class PolicyMatcher:
policies: dict[str, Policy] | None = None,
) -> list[str]:
"""
Filter policies to only those whose conditions match the context.
Filter policies to only those that apply to the given context.
A policy's condition matches if:
- The policy has no condition (condition is None), OR
- The policy's condition evaluates to True for the given context
A policy applies when any policy in its inheritance chain has no
condition or a condition that evaluates to True for the context. The
resolver then drops only the chain members whose own condition fails,
so a child whose condition misses still contributes the guardrails of
its unconditional ancestors. A missing policy resolves to an empty
chain and does not apply.
Args:
policy_names: List of policy names to filter
@ -172,19 +208,11 @@ class PolicyMatcher:
policies: Dictionary of all policies (if None, uses global registry)
Returns:
List of policy names whose conditions match the context
List of policy names that apply to the context
"""
from litellm.proxy.policy_engine.condition_evaluator import ConditionEvaluator
resolved: Final = policies if policies is not None else PolicyMatcher._registry_policies()
matching_policies: Final = []
for policy_name in policy_names:
policy = resolved.get(policy_name)
if policy is None:
continue
# Policy matches if it has no condition OR condition evaluates to True
if policy.condition is None or ConditionEvaluator.evaluate(policy.condition, context):
matching_policies.append(policy_name)
return matching_policies
return [
policy_name
for policy_name in policy_names
if PolicyMatcher._applying_chain_members(policy_name, context, resolved)
]

View file

@ -210,6 +210,7 @@ class PolicyResolver:
Returns:
List of (policy_name, GuardrailPipeline) tuples
"""
from litellm.proxy.policy_engine.condition_evaluator import ConditionEvaluator
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
@ -230,6 +231,11 @@ class PolicyResolver:
policy = policies.get(policy_name)
if policy is None:
continue
if policy.condition is not None and not ConditionEvaluator.evaluate(
condition=policy.condition, context=context
):
verbose_proxy_logger.debug("Policy '%s' condition did not match, skipping pipeline", policy_name)
continue
if policy.pipeline is not None:
pipelines.append((policy_name, policy.pipeline))
verbose_proxy_logger.debug(

View file

@ -13,6 +13,7 @@ import json
from typing import TYPE_CHECKING, Any, Final
from fastapi import APIRouter, Depends, HTTPException
from typing_extensions import ReadOnly, TypedDict
if TYPE_CHECKING:
from prisma.models import LiteLLM_ManagedVectorStoresTable as _VectorStoreRow
@ -56,6 +57,32 @@ def _row_to_vector_store(row: "_VectorStoreRow") -> LiteLLM_ManagedVectorStore:
return LiteLLM_ManagedVectorStore(**row.model_dump())
class _ConfigOwnedDetail(TypedDict):
error: ReadOnly[str]
vector_store_id: ReadOnly[str]
def _raise_if_config_owned(vector_store_id: str) -> None:
if litellm.vector_store_registry is None or not litellm.vector_store_registry.is_config_vector_store(
vector_store_id
):
return
detail: Final[_ConfigOwnedDetail] = {
"error": (
f"Vector store {vector_store_id} is defined in the config file, so the config file owns it and it "
"cannot be changed here. Edit the config file to change it, or remove it from the file to let the "
"database own it."
),
"vector_store_id": vector_store_id,
}
raise HTTPException(status_code=400, detail=detail)
def _with_ownership(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore:
ownership: Final = LiteLLM_ManagedVectorStore(is_config=vector_store.get("is_config", False))
return vector_store | ownership
_LITELLM_PARAMS_MASKER: Final = SensitiveDataMasker(extra_sensitive_patterns=frozenset(("connection",)))
@ -274,6 +301,7 @@ async def new_vector_store(
status_code=400,
detail="vector_store_id and custom_llm_provider are required",
)
_raise_if_config_owned(vector_store_id)
# Extract and validate metadata
metadata: Final = vector_store.get("vector_store_metadata")
@ -306,6 +334,8 @@ async def new_vector_store(
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",
"vector_store": response_vs,
}
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("Error creating vector store: %s", e)
raise HTTPException(status_code=500, detail=str(e))
@ -331,7 +361,9 @@ async def list_vector_stores(
"""
List all available vector stores with optional filtering and pagination.
Combines both in-memory vector stores and those stored in the database.
Database is the source of truth - deleted stores are removed from memory, updated stores sync to memory.
Database is the source of truth for stores it owns: deleted stores are removed from memory, updated stores
sync to memory. Stores declared in the config file are owned by the config file, are always listed, and are
never overwritten by database rows.
Parameters:
- page: int - Page number for pagination (default: 1)
@ -366,8 +398,10 @@ async def list_vector_stores(
if not vector_store_id:
continue
if vector_store.get("is_config", False):
vector_store_map[vector_store_id] = vector_store
# If vector store is in memory but NOT in database, it was deleted
if vector_store_id not in db_vector_store_ids:
elif vector_store_id not in db_vector_store_ids:
verbose_proxy_logger.info(
"Vector store %s exists in memory but not in database - marking for deletion from cache",
vector_store_id,
@ -394,7 +428,7 @@ async def list_vector_stores(
# Filter vector stores based on access control
accessible_vector_stores: Final = []
for vs in await filter_listable_vector_stores(vector_store_map.values(), user_api_key_dict):
redacted = LiteLLM_ManagedVectorStore(**vs)
redacted = _with_ownership(vs)
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
accessible_vector_stores.append(redacted)
@ -467,6 +501,7 @@ async def delete_vector_store(
status_code=404,
detail=f"Vector store with ID {data.vector_store_id} not found",
)
_raise_if_config_owned(data.vector_store_id)
# Check access control
if vector_store_to_check and not await _check_vector_store_access(vector_store_to_check, user_api_key_dict):
@ -545,6 +580,7 @@ async def get_vector_store_info(
litellm_params=_redact_sensitive_litellm_params(vector_store.get("litellm_params")),
team_id=vector_store.get("team_id") or None,
user_id=vector_store.get("user_id") or None,
is_config=vector_store.get("is_config", False),
)
return {"vector_store": vector_store_pydantic_obj}
@ -591,6 +627,7 @@ async def update_vector_store(
update_data: Final = data.model_dump(exclude_unset=True)
vector_store_id: Final[str] = data.vector_store_id
update_data.pop("vector_store_id")
_raise_if_config_owned(vector_store_id)
# Per-store access control: anyone authenticated who passes the
# premium-feature gate could otherwise update *any* vector store —

View file

@ -44,6 +44,8 @@ class LiteLLM_ManagedVectorStore(TypedDict, total=False):
team_id: str | None
user_id: str | None
is_config: ReadOnly[bool]
class LiteLLM_ManagedVectorStoreListResponse(TypedDict, total=False):
"""Response format for listing vector stores"""

View file

@ -340,7 +340,7 @@ class VectorStoreRegistry:
# Verify vector store still exists in database (if we have DB access)
# This ensures deleted vector stores are removed from cache
if vector_store is not None and prisma_client is not None:
if vector_store is not None and prisma_client is not None and not vector_store.get("is_config", False):
try:
# Check if it still exists in database
db_vector_store = await ManagedVectorStoresRepository(prisma_client).table.find_unique(
@ -426,6 +426,7 @@ class VectorStoreRegistry:
vector_store_metadata=vector_store_litellm_params.get("vector_store_metadata"),
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
is_config=True,
)
self.vector_stores.append(litellm_managed_vector_store)
@ -452,6 +453,10 @@ class VectorStoreRegistry:
return response
def is_config_vector_store(self, vector_store_id: str) -> bool:
vector_store: Final = self.get_litellm_managed_vector_store_from_registry(vector_store_id=vector_store_id)
return vector_store is not None and vector_store.get("is_config", False)
def add_vector_store_to_registry(self, vector_store: LiteLLM_ManagedVectorStore):
"""
Add a vector store to the registry
@ -475,10 +480,11 @@ class VectorStoreRegistry:
]
def update_vector_store_in_registry(self, vector_store_id: str, updated_data: LiteLLM_ManagedVectorStore):
"""Update or add a vector store in the registry"""
"""Update or add a vector store in the registry. Config-defined stores are left untouched"""
for i, vector_store in enumerate(self.vector_stores):
if vector_store.get("vector_store_id") == vector_store_id:
self.vector_stores[i] = updated_data
if not vector_store.get("is_config", False):
self.vector_stores[i] = updated_data
return
self.vector_stores.append(updated_data)

File diff suppressed because it is too large Load diff

View file

@ -17,6 +17,7 @@ from models import (
AnthropicMessagesResponse,
ChatBody,
ChatMessage,
ChatMetadata,
ChatResponse,
ChatTool,
KeyGenerateBody,
@ -133,6 +134,31 @@ class GuardrailCreateResponse(BaseModel):
guardrail_id: str
class PolicyConditionBody(BaseModel):
model: str
class PolicyCreateBody(BaseModel):
policy_name: str
inherit: str | None = None
guardrails_add: list[str]
condition: PolicyConditionBody | None = None
class PolicyCreateResponse(BaseModel):
policy_id: str
policy_name: str
class PolicyAttachmentCreateBody(BaseModel):
policy_name: str
tags: list[str]
class PolicyAttachmentCreateResponse(BaseModel):
attachment_id: str
class ApplyGuardrailRequest(BaseModel):
guardrail_name: str
text: str
@ -243,6 +269,49 @@ class GuardrailsClient:
response_type=NoBody,
)
def create_policy(self, body: PolicyCreateBody) -> str:
"""Create a policy via POST /policies and return its name once every replica
can be expected to serve it (policies reach the data plane on the periodic
DB sync, same as guardrails)."""
created = unwrap(
self.proxy.transport.post(
"/policies",
headers=self.proxy.transport.master,
json=body,
response_type=PolicyCreateResponse,
)
)
settle_propagation(time.monotonic())
return created.policy_name
def delete_policy(self, policy_name: str) -> None:
_ = self.proxy.transport.delete(
f"/policies/name/{policy_name}/all-versions",
headers=self.proxy.transport.master,
json=NoBody(),
response_type=NoBody,
)
def attach_policy_to_tags(self, policy_name: str, tags: list[str]) -> str:
attachment_id = unwrap(
self.proxy.transport.post(
"/policies/attachments",
headers=self.proxy.transport.master,
json=PolicyAttachmentCreateBody(policy_name=policy_name, tags=tags),
response_type=PolicyAttachmentCreateResponse,
)
).attachment_id
settle_propagation(time.monotonic())
return attachment_id
def delete_policy_attachment(self, attachment_id: str) -> None:
_ = self.proxy.transport.delete(
f"/policies/attachments/{attachment_id}",
headers=self.proxy.transport.master,
json=NoBody(),
response_type=NoBody,
)
def create_team_opted_out_of_global_guardrails(self, alias: str) -> str:
team_id = unwrap(
self.proxy.transport.post(
@ -322,11 +391,13 @@ class GuardrailsClient:
max_tokens: int = 16,
tools: list[ChatTool] | None = None,
tool_choice: str | None = None,
tags: list[str] | None = None,
) -> StreamingResponse:
"""Drive /chat/completions returning the raw HTTP outcome, for the
assertions a typed body cannot carry: the `x-litellm-applied-guardrails`
response header, which is how an ALLOW scenario proves the guardrail ran
rather than being absent."""
rather than being absent. `tags` land in `metadata.tags`, which is what a
tag-scoped policy attachment matches on."""
return self.proxy.transport.send(
"/chat/completions",
headers=self.proxy.transport.bearer(key),
@ -337,6 +408,7 @@ class GuardrailsClient:
guardrails=guardrails,
tools=tools,
tool_choice=tool_choice,
metadata=ChatMetadata(tags=tags) if tags is not None else None,
),
)

View file

@ -0,0 +1,127 @@
"""Live e2e: a policy attached to a request keeps its inherited parent guardrails
when only the child's own `condition` fails to match the request model.
The parent policy has no condition and adds a content filter. The child inherits
it, adds a second content filter, and carries a model condition. The attachment
points at the child only, so the parent is reachable through inheritance alone.
A request the child condition does not match must still be blocked by the
parent's filter; a request it does match must be blocked by both.
Uses litellm_content_filter (keyword match, no external service) so the block is
deterministic and free, with the request model routed to a real provider.
"""
from __future__ import annotations
import pytest
from e2e_config import CHEAP_OPENAI_MODEL, unique_marker
from e2e_http import StreamingResponse
from guardrails_client import (
GuardrailsClient,
PolicyConditionBody,
PolicyCreateBody,
)
from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
MODEL = CHEAP_OPENAI_MODEL
def _applied_guardrails(outcome: StreamingResponse) -> frozenset[str]:
return frozenset(
name.strip() for name in outcome.headers.get("x-litellm-applied-guardrails", "").split(",") if name.strip()
)
def _setup_child_policy_attached_to_tag(
client: GuardrailsClient,
resources: ResourceManager,
*,
child_condition_model: str,
parent_banned: str,
child_banned: str,
tag: str,
) -> tuple[str, str]:
"""Register parent and child content filters, a parent policy adding the parent
filter, a child policy inheriting it with `child_condition_model`, and attach
only the child to `tag`. Returns (parent_guardrail_name, child_guardrail_name)."""
parent_guardrail = f"e2e-parent-guard-{parent_banned}"
child_guardrail = f"e2e-child-guard-{child_banned}"
parent_guardrail_id = client.create_content_filter_guardrail(parent_guardrail, parent_banned, default_on=False)
resources.defer(lambda: client.delete_guardrail(parent_guardrail_id))
child_guardrail_id = client.create_content_filter_guardrail(child_guardrail, child_banned, default_on=False)
resources.defer(lambda: client.delete_guardrail(child_guardrail_id))
parent_policy = client.create_policy(
PolicyCreateBody(policy_name=f"e2e-parent-policy-{parent_banned}", guardrails_add=[parent_guardrail])
)
resources.defer(lambda: client.delete_policy(parent_policy))
child_policy = client.create_policy(
PolicyCreateBody(
policy_name=f"e2e-child-policy-{child_banned}",
inherit=parent_policy,
guardrails_add=[child_guardrail],
condition=PolicyConditionBody(model=child_condition_model),
)
)
resources.defer(lambda: client.delete_policy(child_policy))
attachment_id = client.attach_policy_to_tags(child_policy, [tag])
resources.defer(lambda: client.delete_policy_attachment(attachment_id))
return parent_guardrail, child_guardrail
class TestPolicyInheritedGuardrail:
def test_child_condition_miss_still_applies_inherited_parent_guardrail(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
parent_banned = unique_marker()
child_banned = unique_marker()
tag = f"e2e-policy-tag-{unique_marker()}"
parent_guardrail, child_guardrail = _setup_child_policy_attached_to_tag(
client,
resources,
child_condition_model=f"never-matches-{unique_marker()}",
parent_banned=parent_banned,
child_banned=child_banned,
tag=tag,
)
outcome = client.chat_raw(scoped_key, MODEL, f"Reply with the single word OK. {parent_banned}", tags=[tag])
assert outcome.status_code == 400, (
f"the inherited parent content filter must block the banned keyword even though the child "
f"policy's own model condition does not match {MODEL}; got {outcome.status_code}: {outcome.body[:300]}"
)
assert parent_guardrail in _applied_guardrails(outcome), (
f"x-litellm-applied-guardrails must name the inherited parent guardrail; got {outcome.headers}"
)
assert child_guardrail not in _applied_guardrails(outcome), (
f"the child's own guardrail must not run when its condition fails; got {outcome.headers}"
)
def test_child_condition_match_applies_child_and_inherited_parent_guardrails(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
parent_banned = unique_marker()
child_banned = unique_marker()
tag = f"e2e-policy-tag-{unique_marker()}"
parent_guardrail, child_guardrail = _setup_child_policy_attached_to_tag(
client,
resources,
child_condition_model=MODEL,
parent_banned=parent_banned,
child_banned=child_banned,
tag=tag,
)
outcome = client.chat_raw(scoped_key, MODEL, f"Reply with the single word OK. {child_banned}", tags=[tag])
assert outcome.status_code == 400, (
f"the child's own content filter must block its banned keyword when the condition matches {MODEL}; "
f"got {outcome.status_code}: {outcome.body[:300]}"
)
assert {parent_guardrail, child_guardrail} <= _applied_guardrails(outcome), (
f"both the child and inherited parent guardrails must run; got {outcome.headers}"
)

View file

@ -12,9 +12,9 @@ The generated lifecycle models use 20 examples, eight steps, generation and shri
Reuse the existing canned provider handlers through `_support/upstream.py`. It rejects internal request fields and exposes actual received requests for independent assertions. Register every created resource for cleanup immediately, keep expected values independent of production calculations, and assert readback plus the runtime effect of a change
The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, skipped tests, failed cleanup or a selected test without a passed call fail qualification. Existing GitHub Actions jobs do not own these tests
The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, failed cleanup or a selected test with neither a passed call nor a skip fail qualification. Skipped nodes are listed under `skipped` in `execution.json`, so the skip reasons double as the open bug list. Existing GitHub Actions jobs do not own these tests
Define integration contract IDs and their canonical test nodes in `contracts.json`. Every node must declare the same IDs with `covers`. The runner checks exact collected and passed selections against that mapping. These IDs belong to this CircleCI suite and must not be added to the separate E2E coverage registry. A manifest declaration alone does not mean a test passed
Define integration contract IDs and their canonical test nodes in `contracts.json`. Every node must declare the same IDs with `covers`. The runner checks exact collected and passed-or-skipped selections against that mapping. These IDs belong to this CircleCI suite and must not be added to the separate E2E coverage registry. A manifest declaration alone does not mean a test passed
Provider sentinels currently use the controlled server, not live recordings. The provider shard also runs the existing strict replay controls for changed requests, exhausted interactions, leftover interactions and no provider connection. Future recorded scenarios must use that replay-only implementation; missing recordings cannot fall back to a real provider. The observation endpoint is destructive and the current selection runs serially against one owned upstream
@ -22,7 +22,7 @@ Fixtures must contain synthetic data only. Keep private incident records and sou
Database cases own their temporary schemas, roles, constraints and proxy processes. They prove reader-versus-writer execution with PostgreSQL lock observations, exercise real transaction wait limits and verify rollback after a reached database failure
Accounting cases compare persisted input and output cost components against literal rates, including zero and default prices. Cache state models assert actual upstream calls, response identity and every persisted charge. Generated accounting tests have a 180-second test limit to accommodate the asynchronous spend writer; CircleCI keeps the whole shard capped at 11 minutes
Accounting cases compare persisted input and output cost components against literal rates, including zero and default prices. Cache state models assert actual upstream calls, response identity and every persisted charge. Generated accounting tests have a 180-second test limit to accommodate the asynchronous spend writer
Provider contracts exercise actual TCP requests with synthetic credentials and local protocol peers. The S3 verifier uses independently implemented equations, a published known-answer vector, a fixed signing clock and deliberately invalid signed requests. Bedrock cases clear ambient AWS credential sources and check the literal model path, loaded role references, STS requests and bearer-only behavior

View file

@ -1,6 +1,6 @@
import os
import socket
import signal
import socket
import subprocess
import sys
import time
@ -13,7 +13,6 @@ from typing import Final
import httpx
import psutil
from integration._support.client import Gateway
@ -61,9 +60,10 @@ def owned_proxy(
*,
config: Path | None = None,
remove_environment: tuple[str, ...] = (),
workers: int = 1,
) -> Iterator[Gateway]:
with owned_proxy_process(
gateway, directory, overrides, config=config, remove_environment=remove_environment
gateway, directory, overrides, config=config, remove_environment=remove_environment, workers=workers
) as owned:
yield owned.gateway
@ -76,11 +76,12 @@ def owned_proxy_process(
*,
config: Path | None = None,
remove_environment: tuple[str, ...] = (),
workers: int = 1,
) -> Iterator[OwnedProxy]:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
port: Final = reserve.getsockname()[1]
root: Final = Path(__file__).resolve().parents[3]
root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
environment: Final = {
**{name: value for name, value in os.environ.items() if name not in remove_environment},
"LITELLM_MASTER_KEY": gateway.key,
@ -104,7 +105,7 @@ def owned_proxy_process(
"--port",
str(port),
"--num_workers",
"1",
str(workers),
"--use_prisma_db_push",
"--enforce_prisma_migration_check",
],

View file

@ -152,6 +152,31 @@ class Provider:
)
return await chat_completions(request)
async def vector_store_search(self, request: Request) -> Response:
body: Final = JSON_OBJECT.validate_json(await request.body())
self.observations.put(Observation(request.url.path, request.headers.get("authorization", ""), body))
query: Final = body.get("query")
if not isinstance(query, str) or not query:
return JSONResponse({"error": {"message": "query is required"}}, status_code=400)
vector_store_id: Final = cast(str, request.path_params["vector_store_id"])
return JSONResponse(
{
"object": "vector_store.search_results.page",
"search_query": query,
"data": [
{
"file_id": f"file_{vector_store_id}",
"filename": "scripted.txt",
"score": 0.9,
"attributes": {},
"content": [{"type": "text", "text": f"scripted context for {query}"}],
}
],
"has_more": False,
"next_page": None,
}
)
async def script(self, request: Request) -> Response:
name: Final = cast(str, request.path_params["model"])
if request.method in {"DELETE", "GET"} and name not in self.scripts:
@ -338,6 +363,7 @@ class Provider:
Route("/v1/completions", completions, methods=["POST"]),
Route("/v1/embeddings", embeddings, methods=["POST"]),
Route("/v1/moderations", moderations, methods=["POST"]),
Route("/vector_stores/{vector_store_id}/search", self.vector_store_search, methods=["POST"]),
Route("/{path:path}", self.scripted, methods=["POST"]),
Route("/{path:path}", self.scripted, methods=["GET"]),
WebSocketRoute("/v1/realtime", self.realtime),

View file

@ -2,11 +2,13 @@ from __future__ import annotations
import ssl
import threading
import time
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
@ -26,6 +28,8 @@ class Reply:
chunks: tuple[bytes, ...] | None = None
abort_after: int | None = None
gate_after_first: threading.Event | None = None
pause_between_chunks: float = 0
headers: Mapping[str, str] = MappingProxyType({})
@dataclass(frozen=True, slots=True)
@ -64,6 +68,8 @@ def wire_server(respond: Callable[[Request], Reply], tls: ssl.SSLContext | None
reply = Reply(status=500)
self.send_response(reply.status)
self.send_header("content-type", reply.content_type)
for name, value in reply.headers.items():
self.send_header(name, value)
if reply.chunks is None:
self.send_header("content-length", str(len(reply.body)))
else:
@ -81,6 +87,8 @@ def wire_server(respond: Callable[[Request], Reply], tls: ssl.SSLContext | None
self.wfile.flush()
if index == 0 and reply.gate_after_first is not None:
assert reply.gate_after_first.wait(timeout=5), "Stream barrier was never released"
if reply.pause_between_chunks and index + 1 < len(reply.chunks):
time.sleep(reply.pause_between_chunks)
else:
self.wfile.write(b"0\r\n\r\n")
self.wfile.flush()

View file

@ -81,17 +81,22 @@ def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
collected: Final = session.config.stash.get(COLLECTED, ())
reports: Final = tuple(report for report in session.config.stash[REPORTS] if report.nodeid in collected)
passed: Final = tuple(report.nodeid for report in reports if report.when == "call" and report.passed)
skipped: Final = tuple(report.nodeid for report in reports if report.skipped)
complete: Final = (
exitstatus == 0
and bool(collected)
and sorted(collected) == sorted(passed)
and all(report.passed for report in reports)
and sorted(collected) == sorted(passed + skipped)
and not any(report.failed for report in reports)
)
output: Final = Path(destination)
output.mkdir(parents=True, exist_ok=True)
(output / "execution.json").write_text(
json.dumps({
"collected": collected, "passed": passed, "complete": complete, "exitstatus": exitstatus,
"collected": collected,
"passed": passed,
"skipped": skipped,
"complete": complete,
"exitstatus": exitstatus,
"hypothesis_version": version("hypothesis"),
"hypothesis_seed": session.config.getoption("hypothesis_seed"),
"order_seed": session.config.getoption("integration_order_seed"),

View file

@ -290,6 +290,12 @@
"tests/integration/spend/test_spend_calculate.py::test_spend_calculate_rejects_unpriced_model_with_400": [
"quota_management.spend_tracking.spend_calculate.rejects_unpriced_model"
],
"tests/integration/spend/test_spend_calculate.py::test_live_preview_entry_charges_cached_tokens_at_the_fresh_rate[gemini-live-2.5-flash-preview-native-audio-09-2025]": [
"quota_management.spend_tracking.spend_calculate.live_preview_cached_tokens_cost_fresh_rate"
],
"tests/integration/spend/test_spend_calculate.py::test_live_preview_entry_charges_cached_tokens_at_the_fresh_rate[gemini/gemini-live-2.5-flash-preview-native-audio-09-2025]": [
"quota_management.spend_tracking.spend_calculate.live_preview_cached_tokens_cost_fresh_rate"
],
"tests/integration/management/test_partial_update_sequences.py::test_restricted_actor_cannot_detach_key_from_project": [
"mgmt.key.update.project_detach_denied_to_restricted_actor"
],
@ -1783,6 +1789,225 @@
"tests/integration/mcp/test_mcp_lifecycle.py::test_same_url_server_grants_scope_discovery_and_direct_or_virtual_execution[bearer]": [
"other.mcp.permissions.same_url_servers_enforce_discovery_and_execution"
],
"tests/integration/management/test_budget_updates.py::test_shortening_budget_duration_moves_reset_at_onto_the_new_schedule": [
"mgmt.budget.update.duration_change_recomputes_reset_at"
],
"tests/integration/management/test_organization_budget_clear.py::test_patch_organization_update_with_null_tpm_limit_clears_it_and_keeps_sibling_limits": [
"mgmt.organization.update.null_clears_budget_limit"
],
"tests/integration/management/test_team_budget_duration_defaults.py::test_team_new_explicit_null_budget_duration_is_not_replaced_by_default": [
"mgmt.team.new.explicit_null_budget_duration_overrides_default"
],
"tests/integration/management/test_team_member_budget_cache.py::test_team_member_default_budget_lands_in_redis_after_first_member_call": [
"mgmt.team_member_budget.default_budget_is_cached_in_redis_as_json"
],
"tests/integration/observability/test_callback_delivery.py::test_streamed_responses_success_callback_carries_provider_apim_request_id": [
"other.observability.callbacks.streamed_responses_events_carry_provider_response_headers"
],
"tests/integration/observability/test_guardrail_effects.py::test_bedrock_passthrough_converse_guardrail_ignores_denied_term_in_tool_definition": [
"other.observability.guardrails.bedrock_passthrough_converse_scans_only_caller_content"
],
"tests/integration/pricing/test_configured_prices.py::test_cost_estimate_reports_configured_prices_for_model_absent_from_cost_map": [
"quota_management.cost_estimate.configured_price.reported_for_model_absent_from_cost_map"
],
"tests/integration/pricing/test_configured_prices.py::test_saving_echoed_model_info_does_not_freeze_cost_map_price_into_deployment": [
"pricing.model_update.echoed_cost_map_price_is_not_persisted_as_override"
],
"tests/integration/pricing/test_databricks_cache_pricing.py::test_databricks_cached_prompt_tokens_bill_at_cache_rates_not_input_rate": [
"pricing.databricks.cached_prompt_tokens_bill_at_cache_rates"
],
"tests/integration/pricing/test_ocr_page_pricing.py::test_ocr_annotation_pages_are_billed_at_annotation_cost_per_page": [
"pricing.ocr.annotation_pages_billed_at_annotation_rate"
],
"tests/integration/pricing/test_service_tier_pricing.py::test_ultrafast_service_tier_bills_ultrafast_rates_and_keeps_pricing_off_the_wire": [
"quota_management.spend_tracking.service_tier_pricing.ultrafast_bills_ultrafast_rates"
],
"tests/integration/spend/test_batch_completion_accounting.py::test_completed_batch_spend_row_records_reasoning_tokens_and_error_file_failures": [
"quota_management.spend_tracking.batch_costs.reasoning_tokens_and_error_file_failures_recorded"
],
"tests/integration/spend/test_batch_observability.py::test_batch_retrieval_row_sums_reasoning_tokens_and_counts_output_and_error_file_failures": [
"spend.batches.retrieval_row_aggregates_reasoning_tokens_and_per_request_counts"
],
"tests/integration/spend/test_batch_poll_starvation.py::test_batches_gone_at_provider_do_not_starve_a_newer_batch_out_of_cost_polling": [
"quota_management.spend_tracking.batch_costs.uncostable_rows_retire_so_newer_batches_are_costed"
],
"tests/integration/spend/test_cache_and_quota.py::test_in_flight_count_tokens_does_not_reserve_key_budget_away_from_a_completion": [
"quota_management.budget.key.in_flight_count_tokens_reserves_nothing_so_completion_reaches_provider"
],
"tests/integration/spend/test_cache_and_quota.py::test_repeated_count_tokens_on_budgeted_key_does_not_reserve_budget_or_block_later_completion": [
"quota_management.budget.key.count_tokens_reserves_nothing_so_completion_within_budget_succeeds"
],
"tests/integration/spend/test_daily_rollup_retry.py::test_failed_daily_user_rollup_commit_is_retried_so_spend_report_and_daily_activity_agree": [
"spend.daily_rollup.failed_user_commit_is_retried_until_report_and_daily_activity_agree"
],
"tests/integration/spend/test_disconnected_bedrock_messages_stream_billing.py::test_client_disconnect_mid_bedrock_messages_stream_still_bills_terminal_usage": [
"spend.anthropic_messages_stream.client_disconnect_bills_terminal_bedrock_usage"
],
"tests/integration/spend/test_failed_dispatch_tokens.py::test_provider_500_after_dispatch_records_estimated_prompt_tokens_on_failure_row": [
"spend.failed_dispatch.failure_row_records_estimated_input_tokens"
],
"tests/integration/spend/test_model_router_selected_model.py::test_model_router_alias_without_router_in_name_keeps_selected_model_in_response_and_spend_log": [
"spend.model_router.selected_model_is_returned_and_persisted_for_plain_alias"
],
"tests/integration/spend/test_org_budget_cli_session_token.py::test_cli_session_token_without_org_id_charges_and_caps_the_team_organization": [
"quota_management.organization_budget.cli_session_token_without_org_id_charges_team_organization"
],
"tests/integration/spend/test_passthrough_budget_reservation.py::test_repeated_gemini_passthrough_calls_stay_served_while_key_spend_is_below_max_budget": [
"spend.budget_reservation.gemini_passthrough_success_releases_reservation_from_spend_counter"
],
"tests/integration/spend/test_team_daily_activity_aggregated.py::test_aggregated_team_activity_reports_the_whole_range_team_spend_in_one_page": [
"quota_management.spend_tracking.team_daily_activity_aggregated_reports_whole_range_team_spend"
],
"tests/integration/spend/test_team_member_spend.py::test_member_added_without_any_budget_is_charged_on_its_membership_row": [
"spend.team_member.member_without_budget_gets_membership_row_and_spend"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_config_store_is_listed_beside_db_store_and_survives_listing": [
"mgmt.vector_store.list.keeps_config_store_beside_db_stores"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_config_store_refuses_new_update_and_delete": [
"mgmt.vector_store.write.config_store_is_read_only"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_db_store_lifecycle_is_unchanged_beside_config_store": [
"mgmt.vector_store.write.db_store_lifecycle_unchanged_beside_config_store"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_chat_with_config_store_searches_upstream_and_injects_context_after_listing": [
"other.vector_store.chat.config_store_search_reaches_upstream_after_listing"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_passthrough_search_on_config_store_uses_yaml_credentials_after_listing": [
"other.vector_store.search.config_store_passthrough_uses_yaml_credentials_after_listing"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_non_admin_key_access_to_config_store_follows_grants_after_admin_listing": [
"authz.vector_store.list.non_admin_key_access_to_config_store_follows_grants"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_peer_process_keeps_config_store_and_sees_db_store_created_elsewhere": [
"mgmt.vector_store.list.peer_process_keeps_config_store_and_sees_db_store"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_concurrent_burst_keeps_config_store_and_refuses_every_config_write": [
"mgmt.vector_store.chaos.concurrent_burst_keeps_config_store_across_workers"
],
"tests/integration/management/test_vector_store_config_ownership.py::test_redis_outage_keeps_config_store_served_and_recovers": [
"mgmt.vector_store.chaos.redis_outage_keeps_config_store_and_recovers"
],
"tests/integration/providers/test_anthropic_advisor_wire.py::test_advisor_sub_call_reaches_the_router_deployment_with_its_key_instead_of_anthropic_unauthenticated": [
"providers.anthropic_messages_advisor.sub_call_uses_the_configured_advisor_deployment"
],
"tests/integration/providers/test_azure_ai_chat_wire.py::test_azure_ai_strips_thinking_blocks_and_cache_control_from_forwarded_messages": [
"providers.azure_ai.anthropic_message_fields_are_stripped_before_foundry"
],
"tests/integration/providers/test_azure_ai_flux2_image_wire.py::test_azure_flux2_flex_generation_hits_flex_provider_path_not_pro": [
"other.provider_wire.azure_ai.flux2_flex_generation_targets_flex_path_with_bfl_body"
],
"tests/integration/providers/test_azure_ai_rerank_auth_wire.py::test_azure_ai_rerank_with_entra_token_and_no_api_key_sends_bearer_to_provider": [
"other.provider_wire.azure_ai.rerank_entra_token_without_api_key_reaches_provider"
],
"tests/integration/providers/test_bedrock_auth_wire.py::test_client_anthropic_oauth_authorization_header_does_not_replace_bedrock_sigv4_signature": [
"providers.bedrock_auth.client_anthropic_oauth_token_never_replaces_sigv4_authorization"
],
"tests/integration/providers/test_bedrock_batch_files_wire.py::test_completions_and_responses_batch_records_upload_as_anthropic_user_messages": [
"other.provider_wire.bedrock.batch_file_completions_and_responses_records_reach_s3_as_user_messages"
],
"tests/integration/providers/test_bedrock_claude_thinking_wire.py::test_prefixed_opus_4_8_reasoning_effort_reaches_bedrock_as_adaptive_thinking_not_budget_tokens": [
"other.provider_wire.bedrock.prefixed_opus_4_8_reasoning_effort_sends_adaptive_thinking"
],
"tests/integration/providers/test_bedrock_converse_config_blocks_wire.py::test_guardrail_and_performance_config_are_not_duplicated_inside_inference_config": [
"other.provider_wire.bedrock.converse_config_blocks_sent_once_at_top_level"
],
"tests/integration/providers/test_bedrock_embedding_wire.py::test_cohere_embed_english_v3_accepts_encoding_format_and_dimensions": [
"other.provider_wire.bedrock.cohere_embed_english_v3_accepts_encoding_format"
],
"tests/integration/providers/test_bedrock_gpt5_reasoning_wire.py::test_gpt5_reasoning_effort_is_accepted_and_sent_as_converse_reasoning_effort": [
"providers.bedrock_converse.gpt5_reasoning_effort_reaches_provider_as_reasoning_effort"
],
"tests/integration/providers/test_bedrock_invoke_tool_search_wire.py::test_gen5_claude_bedrock_invoke_messages_tool_search_sends_bedrock_beta_field": [
"providers.bedrock_invoke.tool_search_gen5_claude_sends_bedrock_beta_and_reports_support"
],
"tests/integration/providers/test_bedrock_mantle_codex_input_wire.py::test_codex_agent_message_context_compaction_and_local_shell_call_reach_mantle_as_supported_items": [
"other.provider_wire.bedrock_mantle.codex_history_items_reach_mantle_as_supported_types"
],
"tests/integration/providers/test_bedrock_mantle_responses_wire.py::test_codex_agent_message_compaction_and_local_shell_items_are_rewritten_for_mantle": [
"providers.bedrock_mantle.codex_history_items_reach_the_wire_as_supported_input_items"
],
"tests/integration/providers/test_bedrock_mantle_wire.py::test_bedrock_mantle_context_overflow_returns_400_saying_prompt_is_too_long": [
"other.provider_wire.bedrock_mantle.context_overflow_is_reported_as_prompt_too_long"
],
"tests/integration/providers/test_bedrock_messages_web_search_replay_wire.py::test_replayed_intercepted_web_search_turn_reaches_bedrock_as_text_and_answers": [
"providers.bedrock_messages.replayed_intercepted_web_search_turn_is_flattened_to_text"
],
"tests/integration/providers/test_bedrock_passthrough_stream_wire.py::test_bedrock_passthrough_converse_stream_response_carries_event_stream_content_type": [
"other.provider_wire.bedrock.passthrough_stream_keeps_event_stream_content_type"
],
"tests/integration/providers/test_bedrock_rerank_wire.py::test_forwarded_client_header_on_rerank_is_excluded_from_the_sigv4_signature": [
"providers.bedrock_rerank.forwarded_client_headers_are_sent_unsigned"
],
"tests/integration/providers/test_bedrock_thinking_tokens_wire.py::test_bedrock_redacted_thinking_is_not_reported_as_zero_reasoning_tokens": [
"other.provider_wire.bedrock.hidden_thinking_tokens_are_not_reported_as_text"
],
"tests/integration/providers/test_dashscope_chat_wire.py::test_dashscope_chat_forwards_reasoning_effort_none_to_the_provider": [
"other.provider_wire.dashscope.reasoning_effort_reaches_provider"
],
"tests/integration/providers/test_databricks_chat_wire.py::test_databricks_stream_final_usage_chunk_reaches_client_and_spend_log": [
"other.provider_wire.databricks.stream_usage_and_cache_reads_reach_client_and_spend_log"
],
"tests/integration/providers/test_databricks_oauth_wire.py::test_databricks_ai_gateway_api_base_requests_oauth_token_from_workspace_origin": [
"other.provider_wire.databricks.oauth_token_url_uses_workspace_origin_for_ai_gateway_api_base"
],
"tests/integration/providers/test_deepseek_vision_wire.py::test_deepseek_vision_forwards_image_url_content_list_instead_of_collapsing_to_text": [
"other.provider_wire.deepseek.vision_image_content_list_reaches_provider"
],
"tests/integration/providers/test_fireworks_ai_router_slug_wire.py::test_fireworks_router_slug_chat_sends_router_resource_not_models_path": [
"other.provider_wire.fireworks_ai.router_slug_chat_sends_router_resource_name"
],
"tests/integration/providers/test_fireworks_ai_router_slug_wire.py::test_fireworks_router_slug_text_completion_sends_router_resource_not_models_path": [
"other.provider_wire.fireworks_ai.router_slug_text_completion_sends_router_resource_name"
],
"tests/integration/providers/test_openai_chat_wire.py::test_openai_chat_tool_choice_without_tools_is_not_forwarded": [
"providers.openai_chat_wire.tool_choice_without_tools_is_dropped_before_the_wire"
],
"tests/integration/providers/test_openai_image_edit_wire.py::test_openai_compatible_image_edit_forwards_seed_form_field_to_backend": [
"other.provider_wire.openai.image_edit_forwards_provider_specific_form_fields"
],
"tests/integration/providers/test_responses_bridge_incomplete.py::test_chat_over_responses_deployment_returns_length_when_output_tokens_run_out": [
"other.provider_wire.responses_bridge.max_output_tokens_incomplete_maps_to_length"
],
"tests/integration/providers/test_tencent_chat_wire.py::test_tencent_thinking_is_sent_in_provider_body_instead_of_failing_the_request[reasoning_effort_none]": [
"other.provider_wire.tencent.thinking_reaches_provider_in_request_body"
],
"tests/integration/providers/test_tencent_chat_wire.py::test_tencent_thinking_is_sent_in_provider_body_instead_of_failing_the_request[thinking_enabled]": [
"other.provider_wire.tencent.thinking_reaches_provider_in_request_body"
],
"tests/integration/providers/test_websearch_interception_wire.py::test_capped_websearch_interception_loop_ends_turn_instead_of_exposing_internal_tool_use": [
"other.provider_wire.anthropic.websearch_interception_capped_loop_ends_turn_without_internal_tool_use"
],
"tests/integration/providers/test_websearch_interception_wire.py::test_streamed_web_search_turn_capped_by_max_agentic_loops_ends_turn_with_snippets_and_ordered_blocks": [
"other.provider_wire.bedrock.websearch_interception_streamed_capped_turn_ends_with_native_results"
],
"tests/integration/providers/test_xai_web_search_wire.py::test_xai_chat_web_search_is_sent_to_responses_with_instructions_and_nested_filters": [
"other.provider_wire.xai.chat_web_search_reaches_responses_with_instructions_and_filters"
],
"tests/integration/routing/test_priority_rate_limit_headers.py::test_non_streaming_v1_messages_success_carries_v3_priority_rate_limit_headers": [
"other.routing.priority_rate_limits.v1_messages_success_exposes_v3_priority_headers"
],
"tests/integration/routing/test_stale_cost_map_boot.py::test_config_deployment_dropped_by_stale_boot_cost_map_is_restored_after_reload": [
"other.routing.cost_map.config_deployment_dropped_by_stale_boot_map_is_restored_after_reload"
],
"tests/integration/streaming/test_stream_contracts.py::test_messages_stream_completes_through_trailing_empty_choices_usage_chunk": [
"other.streaming.messages_bridge.empty_choices_usage_chunk_completes_stream"
],
"tests/integration/streaming/test_stream_contracts.py::test_perplexity_stream_with_cost_breakdown_object_completes_and_bills_total_cost": [
"other.streaming.usage.provider_cost_object_completes_stream_and_bills_total_cost"
],
"tests/integration/streaming/test_stream_contracts.py::test_primary_stream_with_empty_first_chunk_then_disconnect_falls_back_and_bills_the_fallback": [
"other.streaming.fallback.empty_leading_chunk_then_disconnect_streams_fallback_with_usage_and_spend"
],
"tests/integration/streaming/test_stream_contracts.py::test_responses_stream_completes_through_empty_choices_metadata_and_usage_chunks": [
"other.streaming.responses_bridge.empty_choices_chunks_complete_stream"
],
"tests/integration/streaming/test_stream_parallel_slot_release.py::test_failing_stream_logging_callback_does_not_leak_max_parallel_requests_slot": [
"streaming.max_parallel_requests.slot_released_when_stream_logging_callback_fails"
],
"tests/integration/streaming/test_ttft_keepalive.py::test_stream_emits_sse_ping_comments_before_the_first_data_frame_while_upstream_is_silent": [
"streaming.keepalive.sse_pings_fill_silent_time_to_first_token"
],
"tests/integration/observability/test_callback_delivery.py::test_response_survives_raising_success_deployment_hook[chat]": [
"other.observability.callbacks.raising_success_deployment_hook_keeps_response"
],
@ -1795,6 +2020,83 @@
"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",
"mgmt.customer.update.budget_change_returns_promptly_with_wedged_coordination_redis",
"mgmt.key.reset_spend.returns_promptly_with_wedged_coordination_redis",
"mgmt.auth_cache_invalidation.publish_parked_by_short_redis_wedge_lands_after_recovery",
"mgmt.auth_cache_invalidation.burst_with_worker_kill_keeps_serving_while_redis_wedged"
],
"tests/integration/spend/test_team_member_budget_alerts.py::test_team_member_budget_thresholds_email_member_and_configured_recipients": [
"quota_management.budget.team_member.alerts_at_configured_thresholds"
]

View file

@ -0,0 +1,12 @@
model_list: []
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
database_url: os.environ/DATABASE_URL
store_model_in_db: true
disable_spend_logs: false
proxy_batch_write_at: 1
coordination_redis:
host: os.environ/REDIS_HOST
port: os.environ/REDIS_PORT
router_settings:
disable_cooldowns: true

View file

@ -0,0 +1,29 @@
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from tests.integration._support.client import Gateway, string_value
from tests.integration._support.database import read_rows
def _persisted_reset_at(budget_id: str) -> datetime:
rows: Final = read_rows(
'SELECT budget_reset_at::text AS reset_at FROM "LiteLLM_BudgetTable" WHERE budget_id = %s', (budget_id,)
)
assert len(rows) == 1, rows
reset_at: Final = datetime.fromisoformat(string_value(rows[0]["reset_at"]))
return reset_at if reset_at.tzinfo is not None else reset_at.replace(tzinfo=timezone.utc)
@pytest.mark.covers("mgmt.budget.update.duration_change_recomputes_reset_at")
def test_shortening_budget_duration_moves_reset_at_onto_the_new_schedule(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
budget_id: Final = scenario.budget(max_budget=10.0, budget_duration="10d")
ten_day_reset_at: Final = _persisted_reset_at(budget_id)
before: Final = datetime.now(timezone.utc)
response: Final = gateway.request("POST", "/budget/update", {"budget_id": budget_id, "budget_duration": "1d"})
assert response.status_code == 200, response.text
updated: Final = _persisted_reset_at(budget_id)
assert updated < ten_day_reset_at, f"{updated} not before {ten_day_reset_at}"
assert before < updated <= before + timedelta(days=1, minutes=5), f"{updated} not within 1d of {before}"

View file

@ -0,0 +1,59 @@
import uuid
from typing import Final
import pytest
from tests.integration._support.client import Gateway, object_value, string_value
from tests.integration._support.database import read_rows
def _budget_rows(budget_id: str) -> list[dict[str, object]]:
return read_rows(
'SELECT tpm_limit, rpm_limit, max_budget FROM "LiteLLM_BudgetTable" WHERE budget_id = %s', (budget_id,)
)
@pytest.mark.covers("mgmt.organization.update.null_clears_budget_limit")
def test_patch_organization_update_with_null_tpm_limit_clears_it_and_keeps_sibling_limits(gateway: Gateway) -> None:
created: Final = gateway.post(
"/organization/new",
{
"organization_alias": f"integration-{uuid.uuid4().hex}",
"tpm_limit": 4000,
"rpm_limit": 40,
"max_budget": 12.5,
},
)
organization_id: Final = string_value(created["organization_id"])
budget_id: Final = string_value(created["budget_id"])
try:
assert _budget_rows(budget_id) == [{"tpm_limit": 4000, "rpm_limit": 40, "max_budget": 12.5}]
updated: Final = gateway.request(
"PATCH", "/organization/update", {"organization_id": organization_id, "tpm_limit": None}
)
assert updated.status_code == 200, updated.text
updated_budget: Final = object_value(object_value(updated.json())["litellm_budget_table"])
assert (updated_budget["tpm_limit"], updated_budget["rpm_limit"], updated_budget["max_budget"]) == (
None,
40,
12.5,
), updated.text
assert _budget_rows(budget_id) == [{"tpm_limit": None, "rpm_limit": 40, "max_budget": 12.5}]
info: Final = gateway.request("GET", "/organization/info", params={"organization_id": organization_id})
assert info.status_code == 200, info.text
info_budget: Final = object_value(object_value(info.json())["litellm_budget_table"])
assert (info_budget["tpm_limit"], info_budget["rpm_limit"], info_budget["max_budget"]) == (
None,
40,
12.5,
), info.text
finally:
deleted: Final = gateway.request("DELETE", "/organization/delete", {"organization_ids": [organization_id]})
assert deleted.status_code == 200, deleted.text
gateway.post("/budget/delete", {"id": budget_id})
assert (
read_rows(
'SELECT organization_id FROM "LiteLLM_OrganizationTable" WHERE organization_id = %s', (organization_id,)
)
== []
)

View file

@ -0,0 +1,61 @@
import uuid
from pathlib import Path
from typing import Final
import pytest
import yaml
from pydantic import JsonValue
from tests.integration._support.client import Gateway, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.process import owned_proxy
def _budget_row(team_id: str) -> dict[str, JsonValue]:
rows: Final = read_rows(
'SELECT max_budget, budget_duration, budget_reset_at::text FROM "LiteLLM_TeamTable" WHERE team_id = %s',
(team_id,),
)
assert len(rows) == 1, rows
return rows[0]
@pytest.mark.covers("mgmt.team.new.explicit_null_budget_duration_overrides_default")
def test_team_new_explicit_null_budget_duration_is_not_replaced_by_default(gateway: Gateway, tmp_path: Path) -> None:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"]["default_team_params"] = {"budget_duration": "30d"}
path: Final = tmp_path / "team-defaults.yaml"
path.write_text(yaml.safe_dump(config))
with (
owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=path) as candidate,
candidate.scenario() as scenario,
):
never_resetting: Final = candidate.request(
"POST",
"/team/new",
{"team_alias": f"integration-{uuid.uuid4().hex}", "max_budget": 500, "budget_duration": None},
)
assert never_resetting.status_code == 200, never_resetting.text
never_resetting_id: Final = string_value(never_resetting.json()["team_id"])
scenario.cleanups.callback(scenario.delete_team, never_resetting_id)
assert never_resetting.json()["max_budget"] == 500.0, never_resetting.text
assert never_resetting.json()["budget_duration"] is None, never_resetting.text
assert never_resetting.json()["budget_reset_at"] is None, never_resetting.text
assert _budget_row(never_resetting_id) == {
"max_budget": 500.0,
"budget_duration": None,
"budget_reset_at": None,
}
inheriting: Final = candidate.request(
"POST", "/team/new", {"team_alias": f"integration-{uuid.uuid4().hex}", "max_budget": 500}
)
assert inheriting.status_code == 200, inheriting.text
inheriting_id: Final = string_value(inheriting.json()["team_id"])
scenario.cleanups.callback(scenario.delete_team, inheriting_id)
assert inheriting.json()["budget_duration"] == "30d", inheriting.text
assert inheriting.json()["budget_reset_at"] is not None, inheriting.text
inheriting_row: Final = _budget_row(inheriting_id)
assert inheriting_row["max_budget"] == 500.0, inheriting_row
assert inheriting_row["budget_duration"] == "30d", inheriting_row
assert inheriting_row["budget_reset_at"] is not None, inheriting_row

View file

@ -0,0 +1,40 @@
import os
from typing import Final
import pytest
from pydantic import JsonValue, TypeAdapter
from redis import Redis
from tests.integration._support.client import Gateway, eventually, object_value, string_value
from tests.integration._support.database import read_rows
_CACHED_BUDGET: Final = TypeAdapter(dict[str, JsonValue])
@pytest.mark.covers("mgmt.team_member_budget.default_budget_is_cached_in_redis_as_json")
def test_team_member_default_budget_lands_in_redis_after_first_member_call(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
user: Final = scenario.user()
team: Final = scenario.team(team_member_budget=25)
key: Final = scenario.key(team_id=team, user_id=user, models=[model])
teams: Final = read_rows('SELECT metadata FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,))
assert len(teams) == 1, teams
budget_id: Final = string_value(object_value(teams[0]["metadata"])["team_member_budget_id"])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "member budget cache"}]},
key=key,
)
assert response.status_code == 200, response.text
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache:
cached: Final = eventually(
lambda: cache.get(f"team_member_default_budget:{budget_id}"),
lambda value: value is not None,
seconds=10,
)
assert isinstance(cached, bytes), cached
budget: Final = _CACHED_BUDGET.validate_json(cached)
assert budget["budget_id"] == budget_id, cached
assert budget["max_budget"] == 25, cached

View file

@ -0,0 +1,379 @@
import os
import signal
import time
import uuid
from collections.abc import Callable, Mapping
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import Final
from urllib.parse import urlsplit, urlunsplit
import psutil
import psycopg
import pytest
from psycopg import sql
from pydantic import JsonValue
from redis import Redis
from redis.client import PubSub
from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from tests.integration._support.process import owned_proxy
from tests.integration._support.redis_process import owned_redis
_USERS: Final = 60
_BURST: Final = 30
_HANDLER_BUDGET_SECONDS: Final = 0.75
_BULK_BUDGET_SECONDS: Final = 2.0
_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation"
def _timed_post(candidate: Gateway, path: str, body: Mapping[str, JsonValue], timeout: float = 15) -> float:
started: Final = time.monotonic()
response: Final = candidate.client.request(
"POST",
path,
json=body,
headers={"Authorization": f"Bearer {candidate.key}"},
timeout=timeout,
)
elapsed: Final = time.monotonic() - started
assert response.status_code == 200, f"POST {path}: {response.status_code} {response.text} after {elapsed:.3f}s"
return elapsed
def _received(pubsub: PubSub) -> tuple[dict[str, JsonValue], ...]:
messages: list[dict[str, JsonValue]] = []
while True:
message = pubsub.get_message(ignore_subscribe_messages=True, timeout=0)
if message is None:
return tuple(messages)
data = message.get("data")
if isinstance(data, (bytes, str)):
messages.append(JSON_OBJECT.validate_json(data))
def _worker_pid(port: int) -> int:
for process in psutil.process_iter():
parent = process.parent()
if parent is None:
continue
try:
cmdline = parent.cmdline()
own_cmdline = process.cmdline()
except (psutil.NoSuchProcess, psutil.AccessDenied):
continue
if (
"integration._support.proxy" in cmdline
and "--port" in cmdline
and str(port) in cmdline
and not any("prisma" in part for part in own_cmdline)
):
return process.pid
raise AssertionError(f"no uvicorn worker found under the owned proxy on port {port}")
def _burst_call(
index: int, users: tuple[str, ...], key: str, team_id: str, customer_id: str
) -> tuple[str, dict[str, JsonValue]]:
match index % 5:
case 0:
return "/user/update", {"user_id": users[index], "max_budget": 200.0 + index}
case 1:
return "/user/update", {"user_id": users[index], "tpm_limit": 1000 + index}
case 2:
return "/key/update", {"key": key, "max_budget": 7.0 + index}
case 3:
return "/team/update", {"team_id": team_id, "max_budget": 7.0 + index}
case _:
return "/customer/update", {"user_id": customer_id, "max_budget": 7.0 + index}
@pytest.mark.timeout(240)
@pytest.mark.covers(
"mgmt.user.update.budget_change_returns_promptly_with_wedged_coordination_redis",
"mgmt.user.bulk_update.budget_change_returns_promptly_with_wedged_coordination_redis",
"mgmt.customer.update.budget_change_returns_promptly_with_wedged_coordination_redis",
"mgmt.key.reset_spend.returns_promptly_with_wedged_coordination_redis",
"mgmt.auth_cache_invalidation.publish_parked_by_short_redis_wedge_lands_after_recovery",
"mgmt.auth_cache_invalidation.burst_with_worker_kill_keeps_serving_while_redis_wedged",
)
def test_user_budget_updates_return_promptly_while_coordination_redis_is_wedged(
gateway: Gateway, tmp_path: Path, record_property: Callable[[str, object], None]
) -> None:
original: Final = os.environ["DATABASE_URL"]
identity: Final = "integration_wedged_redis_" + uuid.uuid4().hex
parsed: Final = urlsplit(original)
database_url: Final = urlunsplit((parsed.scheme, parsed.netloc, "/" + identity, "", ""))
timings: dict[str, float] = {}
with psycopg.connect(original, autocommit=True) as admin:
admin.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(identity)))
try:
results_dir: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(tmp_path)))
prior_logs: Final = frozenset(results_dir.glob("owned-proxy-*.log"))
with (
owned_redis(tmp_path) as coordination,
owned_proxy(
gateway,
tmp_path,
{
"DATABASE_URL": database_url,
"REDIS_HOST": coordination.host,
"REDIS_PORT": str(coordination.port),
},
config=Path("tests/integration/coordination_redis_proxy_config.yaml"),
workers=2,
) as candidate,
Redis(host=coordination.host, port=coordination.port, socket_timeout=1) as subscriber_client,
):
pubsub: Final = subscriber_client.pubsub()
pubsub.subscribe(_CHANNEL)
received: list[dict[str, JsonValue]] = []
def drained() -> tuple[dict[str, JsonValue], ...]:
received.extend(_received(pubsub))
return tuple(received)
eventually(
lambda: subscriber_client.pubsub_numsub(_CHANNEL)[0][1],
lambda count: count >= 3,
seconds=15,
)
users: Final = tuple(f"{identity}_u{index}" for index in range(_USERS))
for user_id in users:
candidate.post("/user/new", {"user_id": user_id, "auto_create_key": False, "max_budget": 10.0})
key: Final = string_value(
candidate.post("/key/generate", {"user_id": users[0], "max_budget": 5.0})["key"]
)
team_id: Final = string_value(
candidate.post("/team/new", {"team_alias": identity, "max_budget": 5.0})["team_id"]
)
customer_id: Final = identity + "_cust"
candidate.post("/customer/new", {"user_id": customer_id, "max_budget": 5.0})
drained()
timings["h1_healthy"] = _timed_post(
candidate, "/user/update", {"user_id": users[0], "max_budget": 11.0}
)
assert timings["h1_healthy"] < _HANDLER_BUDGET_SECONDS, (
f"healthy /user/update took {timings['h1_healthy']:.3f}s"
)
eventually(
drained,
lambda messages: any(message.get("cache_key") == users[0] for message in messages),
seconds=10,
)
timings["h2_healthy_control"] = _timed_post(
candidate, "/user/update", {"user_id": users[0], "tpm_limit": 1000}
)
assert timings["h2_healthy_control"] < _HANDLER_BUDGET_SECONDS, (
f"healthy control update took {timings['h2_healthy_control']:.3f}s"
)
coordination.signal(signal.SIGSTOP)
try:
timings["s2_wedged_control"] = _timed_post(
candidate, "/user/update", {"user_id": users[0], "tpm_limit": 1000}
)
assert timings["s2_wedged_control"] < _HANDLER_BUDGET_SECONDS, (
f"control update without a cache-relevant field took {timings['s2_wedged_control']:.3f}s"
)
timings["s1_user_update"] = _timed_post(
candidate, "/user/update", {"user_id": users[0], "max_budget": 98.0}
)
assert timings["s1_user_update"] < _HANDLER_BUDGET_SECONDS, (
f"/user/update with max_budget took {timings['s1_user_update']:.3f}s "
"with a wedged coordination Redis"
)
timings["s3_bulk_update"] = _timed_post(
candidate, "/user/bulk_update", {"all_users": True, "user_updates": {"max_budget": 79.0}}
)
assert timings["s3_bulk_update"] < _BULK_BUDGET_SECONDS, (
f"/user/bulk_update over {_USERS} users took {timings['s3_bulk_update']:.3f}s "
"with a wedged coordination Redis"
)
timings["s4_key_update"] = _timed_post(
candidate, "/key/update", {"key": key, "max_budget": 6.0}, timeout=60
)
assert timings["s4_key_update"] < 30, (
f"/key/update hung for {timings['s4_key_update']:.3f}s with a wedged coordination Redis"
)
timings["s5_team_update"] = _timed_post(
candidate, "/team/update", {"team_id": team_id, "max_budget": 6.0}, timeout=60
)
assert timings["s5_team_update"] < 30, (
f"/team/update hung for {timings['s5_team_update']:.3f}s with a wedged coordination Redis"
)
timings["s6_customer_update"] = _timed_post(
candidate, "/customer/update", {"user_id": customer_id, "max_budget": 6.0}
)
assert timings["s6_customer_update"] < _HANDLER_BUDGET_SECONDS, (
f"/customer/update took {timings['s6_customer_update']:.3f}s with a wedged coordination Redis"
)
timings["s7_reset_spend"] = _timed_post(candidate, f"/key/{key}/reset_spend", {"reset_to": 0})
assert timings["s7_reset_spend"] < _HANDLER_BUDGET_SECONDS, (
f"/key/<key>/reset_spend took {timings['s7_reset_spend']:.3f}s with a wedged coordination Redis"
)
missing_started: Final = time.monotonic()
missing: Final = candidate.request(
"POST", "/user/update", {"user_id": users[0], "max_budget": "not-a-number"}
)
timings["s8_invalid_body"] = time.monotonic() - missing_started
assert missing.status_code // 100 == 4, (
f"/user/update with an invalid body returned {missing.status_code} "
f"in {timings['s8_invalid_body']:.3f}s"
)
assert timings["s8_invalid_body"] < _HANDLER_BUDGET_SECONDS, (
f"/user/update with an invalid body took {timings['s8_invalid_body']:.3f}s"
)
def burst_request(path: str, body: Mapping[str, JsonValue]) -> tuple[object, float]:
started: Final = time.monotonic()
try:
response: Final = candidate.client.request(
"POST",
path,
json=body,
headers={"Authorization": f"Bearer {candidate.key}"},
timeout=60,
)
return response.status_code, time.monotonic() - started
except Exception as error: # noqa: BLE001 # the killed worker drops in-flight requests
return error, time.monotonic() - started
port: Final = candidate.client.base_url.port
assert port is not None, f"owned proxy client has no port: {candidate.client.base_url}"
with ThreadPoolExecutor(_BURST) as pool:
futures: Final = [
pool.submit(
burst_request,
*_burst_call(i, users, key, team_id, customer_id),
)
for i in range(_BURST)
]
os.kill(_worker_pid(port), signal.SIGKILL)
results: Final = [future.result() for future in futures]
responses: Final = [(status, elapsed) for status, elapsed in results if isinstance(status, int)]
failures: Final = [status for status, _elapsed in responses if status != 200]
assert not failures, f"burst responses that were not 200: {failures}"
transport_errors: Final = [status for status, _elapsed in results if not isinstance(status, int)]
assert len(transport_errors) <= 3, (
f"{len(transport_errors)} requests raised transport errors: {transport_errors!r}"
)
elapsed_sorted: Final = sorted(
elapsed for i, (status, elapsed) in enumerate(results) if i % 5 in (0, 1, 4) and status == 200
)
timings["c1_burst_p95"] = elapsed_sorted[int(len(elapsed_sorted) * 0.95) - 1]
assert timings["c1_burst_p95"] < _HANDLER_BUDGET_SECONDS, (
f"burst p95 {timings['c1_burst_p95']:.3f}s"
)
eventually(
lambda: candidate.request("GET", "/health/liveliness").status_code,
lambda status: status == 200,
seconds=15,
)
timings["c1_survivor"] = _timed_post(
candidate, "/user/update", {"user_id": users[0], "tpm_limit": 2000}
)
assert timings["c1_survivor"] < _HANDLER_BUDGET_SECONDS, (
f"control update on the surviving worker took {timings['c1_survivor']:.3f}s"
)
finally:
coordination.signal(signal.SIGCONT)
wedged_keys: Final = {users[i] for i in range(_BURST) if i % 5 == 0 and i != 0} | {f"team_id:{team_id}"}
def proxy_log() -> str:
return "".join(
path.read_text() for path in results_dir.glob("owned-proxy-*.log") if path not in prior_logs
)
team_wedged_key: Final = f"team_id:{team_id}"
eventually(
proxy_log,
lambda text: (
all(
f"publish for {wedged_key} failed" in text
for wedged_key in wedged_keys
if wedged_key != team_wedged_key
)
and (
f"publish for {team_wedged_key} failed" in text
or f"internal usage cache entry {team_wedged_key}" in text
)
),
seconds=45,
)
marker: Final = len(received)
drained()
recovered_keys: Final = {str(message.get("cache_key")) for message in received[marker:]}
assert recovered_keys.isdisjoint(wedged_keys), (
f"wedged publishes unexpectedly landed after recovery: {sorted(recovered_keys & wedged_keys)}"
)
coordination.signal(signal.SIGSTOP)
try:
timings["r1b_short_wedge_a"] = _timed_post(
candidate, "/user/update", {"user_id": users[4], "max_budget": 15.0}
)
assert timings["r1b_short_wedge_a"] < _HANDLER_BUDGET_SECONDS, (
f"/user/update inside a short wedge took {timings['r1b_short_wedge_a']:.3f}s"
)
timings["r1b_short_wedge_b"] = _timed_post(
candidate, "/user/update", {"user_id": users[5], "max_budget": 16.0}
)
assert timings["r1b_short_wedge_b"] < _HANDLER_BUDGET_SECONDS, (
f"/user/update inside a short wedge took {timings['r1b_short_wedge_b']:.3f}s"
)
finally:
coordination.signal(signal.SIGCONT)
eventually(
drained,
lambda messages: {str(message.get("cache_key")) for message in messages} >= {users[4], users[5]},
seconds=10,
)
timings["r2_resumed"] = _timed_post(
candidate, "/user/update", {"user_id": users[1], "max_budget": 12.0}
)
assert timings["r2_resumed"] < _HANDLER_BUDGET_SECONDS, (
f"post-recovery /user/update took {timings['r2_resumed']:.3f}s"
)
eventually(
drained,
lambda messages: any(message.get("cache_key") == users[1] for message in messages),
seconds=10,
)
coordination.stop()
timings["f1_refused"] = _timed_post(
candidate, "/user/update", {"user_id": users[2], "max_budget": 13.0}
)
assert timings["f1_refused"] < _HANDLER_BUDGET_SECONDS, (
f"/user/update with refused coordination Redis took {timings['f1_refused']:.3f}s"
)
coordination.start()
restarted_pubsub: Final = subscriber_client.pubsub()
restarted_pubsub.subscribe(_CHANNEL)
restarted_received: list[dict[str, JsonValue]] = []
def drained_after_restart() -> tuple[dict[str, JsonValue], ...]:
restarted_received.extend(_received(restarted_pubsub))
return tuple(restarted_received)
eventually(
lambda: subscriber_client.pubsub_numsub(_CHANNEL)[0][1],
lambda count: count >= 3,
seconds=30,
)
timings["f2_restarted"] = _timed_post(
candidate, "/user/update", {"user_id": users[3], "max_budget": 14.0}
)
assert timings["f2_restarted"] < _HANDLER_BUDGET_SECONDS, (
f"/user/update after Redis restart took {timings['f2_restarted']:.3f}s"
)
eventually(
drained_after_restart,
lambda messages: any(message.get("cache_key") == users[3] for message in messages),
seconds=10,
)
info_last: Final = object_value(candidate.get("/user/info", {"user_id": users[-1]})["user_info"])
assert info_last["max_budget"] == 79.0, info_last
info_user3: Final = object_value(candidate.get("/user/info", {"user_id": users[3]})["user_info"])
assert info_user3["max_budget"] == 14.0, info_user3
finally:
admin.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(identity)))
record_property("cell_elapsed_seconds", timings)

View file

@ -0,0 +1,351 @@
import os
import uuid
from collections.abc import Mapping
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import Final
from urllib.parse import urlsplit, urlunsplit
import httpx
import psycopg
import pytest
from psycopg import sql
from pydantic import JsonValue
from tests.integration._support.client import Gateway, eventually, object_value
from tests.integration._support.database import read_rows
from tests.integration._support.process import owned_proxy
from tests.integration._support.redis_process import owned_redis
CONFIG_STORE_ID: Final = "vs_integration_config_store"
CONFIG_STORE_NAME: Final = "integration-config-store"
SEARCH_PATH: Final = f"/vector_stores/{CONFIG_STORE_ID}/search"
PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml"
def listed_rows(response: httpx.Response) -> tuple[dict[str, JsonValue], ...]:
rows: Final = object_value(response.json()).get("data")
assert isinstance(rows, list), response.text
return tuple(object_value(row) for row in rows)
def listed_store(gateway: Gateway, vector_store_id: str, *, key: str | None = None) -> dict[str, JsonValue]:
listed: Final = gateway.request("GET", "/vector_store/list", key=key)
assert listed.status_code == 200, listed.text
matches: Final = tuple(row for row in listed_rows(listed) if row["vector_store_id"] == vector_store_id)
assert len(matches) == 1, f"{vector_store_id} appears {len(matches)} times in {listed.text}"
return matches[0]
def listed_ids(gateway: Gateway) -> tuple[str, ...]:
rows: Final = gateway.get("/vector_store/list")["data"]
assert isinstance(rows, list)
return tuple(str(object_value(row)["vector_store_id"]) for row in rows)
def config_store_info(gateway: Gateway) -> dict[str, JsonValue]:
return object_value(gateway.post("/vector_store/info", {"vector_store_id": CONFIG_STORE_ID})["vector_store"])
def store_rows(vector_store_id: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT vector_store_id, vector_store_name FROM "LiteLLM_ManagedVectorStoresTable" WHERE vector_store_id = %s',
(vector_store_id,),
)
def assert_config_write_refused(gateway: Gateway) -> None:
for path, body in (
("/vector_store/update", {"vector_store_id": CONFIG_STORE_ID, "vector_store_name": "renamed"}),
("/vector_store/delete", {"vector_store_id": CONFIG_STORE_ID}),
("/vector_store/new", {"vector_store_id": CONFIG_STORE_ID, "custom_llm_provider": "openai"}),
):
refused = gateway.request("POST", path, body)
assert refused.status_code == 400, f"{path}: {refused.status_code} {refused.text}"
error = object_value(object_value(refused.json())["detail"])
assert error["vector_store_id"] == CONFIG_STORE_ID, refused.text
assert "config file" in str(error["error"]), refused.text
def burst_list(gateway: Gateway) -> tuple[int, str]:
response: Final = gateway.request("GET", "/vector_store/list")
if response.status_code != 200:
return response.status_code, response.text
ids: Final = tuple(str(row["vector_store_id"]) for row in listed_rows(response))
return response.status_code, "config" if CONFIG_STORE_ID in ids else response.text
def burst_post(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> tuple[int, str]:
response: Final = gateway.request("POST", path, body)
return response.status_code, response.text
def upstream_requests(upstream: httpx.Client, marker: str) -> list[dict[str, JsonValue]]:
observed: Final = upstream.get("/__observations")
observed.raise_for_status()
requests: Final = object_value(observed.json())["requests"]
assert isinstance(requests, list), observed.text
return [object_value(value) for value in requests if marker in str(object_value(value)["body"])]
@pytest.mark.covers("mgmt.vector_store.list.keeps_config_store_beside_db_stores")
def test_config_store_is_listed_beside_db_store_and_survives_listing(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
db_store_id: Final = f"vs_db_{uuid.uuid4().hex}"
gateway.post("/vector_store/new", {"vector_store_id": db_store_id, "custom_llm_provider": "openai"})
scenario.cleanups.callback(gateway.post, "/vector_store/delete", {"vector_store_id": db_store_id})
before: Final = config_store_info(gateway)
assert before["vector_store_id"] == CONFIG_STORE_ID, before
config_row: Final = listed_store(gateway, CONFIG_STORE_ID)
assert config_row["is_config"] is True, config_row
assert config_row["vector_store_name"] == CONFIG_STORE_NAME, config_row
assert object_value(config_row["litellm_params"])["api_key"] != "integration-provider-key", config_row
db_row: Final = listed_store(gateway, db_store_id)
assert db_row["is_config"] is False, db_row
after: Final = config_store_info(gateway)
assert after["vector_store_id"] == CONFIG_STORE_ID, after
assert after["is_config"] is True, after
assert after["vector_store_description"] == "declared in tests/integration/proxy_config.yaml", after
assert store_rows(CONFIG_STORE_ID) == [], "config store must not need a database row"
assert listed_store(gateway, CONFIG_STORE_ID)["is_config"] is True
@pytest.mark.covers("mgmt.vector_store.write.config_store_is_read_only")
def test_config_store_refuses_new_update_and_delete(gateway: Gateway) -> None:
assert_config_write_refused(gateway)
row: Final = listed_store(gateway, CONFIG_STORE_ID)
assert row["vector_store_name"] == CONFIG_STORE_NAME, row
assert row["is_config"] is True, row
assert config_store_info(gateway)["vector_store_name"] == CONFIG_STORE_NAME
@pytest.mark.covers("mgmt.vector_store.write.db_store_lifecycle_unchanged_beside_config_store")
def test_db_store_lifecycle_is_unchanged_beside_config_store(gateway: Gateway) -> None:
incomplete: Final = gateway.request("POST", "/vector_store/new", {"custom_llm_provider": "openai"})
assert incomplete.status_code == 400, incomplete.text
db_store_id: Final = f"vs_db_{uuid.uuid4().hex}"
created: Final = gateway.request(
"POST",
"/vector_store/new",
{"vector_store_id": db_store_id, "custom_llm_provider": "openai", "vector_store_name": "first"},
)
assert created.status_code == 200, created.text
assert store_rows(db_store_id) == [{"vector_store_id": db_store_id, "vector_store_name": "first"}]
updated: Final = gateway.post(
"/vector_store/update", {"vector_store_id": db_store_id, "vector_store_name": "second"}
)
assert object_value(updated["vector_store"])["vector_store_name"] == "second", updated
assert store_rows(db_store_id) == [{"vector_store_id": db_store_id, "vector_store_name": "second"}]
row: Final = listed_store(gateway, db_store_id)
assert row["vector_store_name"] == "second" and row["is_config"] is False, row
info: Final = object_value(gateway.post("/vector_store/info", {"vector_store_id": db_store_id})["vector_store"])
assert info["vector_store_name"] == "second" and info["is_config"] is False, info
gateway.post("/vector_store/delete", {"vector_store_id": db_store_id})
assert store_rows(db_store_id) == []
assert db_store_id not in listed_ids(gateway)
assert CONFIG_STORE_ID in listed_ids(gateway)
missing: Final = gateway.request("POST", "/vector_store/info", {"vector_store_id": db_store_id})
assert missing.status_code == 404, missing.text
@pytest.mark.covers("other.vector_store.chat.config_store_search_reaches_upstream_after_listing")
def test_chat_with_config_store_searches_upstream_and_injects_context_after_listing(gateway: Gateway) -> None:
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model()
marker: Final = f"lit6337 {uuid.uuid4().hex}"
assert CONFIG_STORE_ID in listed_ids(gateway)
upstream.get("/__observations").raise_for_status()
completion: Final = gateway.post(
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": marker}], "vector_store_ids": [CONFIG_STORE_ID]},
)
assert object_value(completion["usage"])["total_tokens"] == 40, completion
requests: Final = upstream_requests(upstream, marker)
searches: Final = [value for value in requests if value["path"] == SEARCH_PATH]
assert len(searches) == 1, requests
assert object_value(searches[0]["body"])["query"] == marker, searches
assert searches[0]["authorization"] == "Bearer integration-provider-key", searches
chats: Final = [value for value in requests if value["path"] == "/v1/chat/completions"]
assert len(chats) == 1, requests
messages: Final = object_value(chats[0]["body"])["messages"]
assert isinstance(messages, list), chats
contents: Final = tuple(str(object_value(message)["content"]) for message in messages)
assert contents == (f"Context:\n\nscripted context for {marker}\n\n", marker), contents
@pytest.mark.covers("other.vector_store.search.config_store_passthrough_uses_yaml_credentials_after_listing")
def test_passthrough_search_on_config_store_uses_yaml_credentials_after_listing(gateway: Gateway) -> None:
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
marker: Final = f"lit6337 passthrough {uuid.uuid4().hex}"
assert CONFIG_STORE_ID in listed_ids(gateway)
upstream.get("/__observations").raise_for_status()
searched: Final = gateway.request("POST", f"/v1/vector_stores/{CONFIG_STORE_ID}/search", {"query": marker})
assert searched.status_code == 200, searched.text
data: Final = listed_rows(searched)
assert len(data) == 1, searched.text
content: Final = data[0]["content"]
assert isinstance(content, list), searched.text
assert object_value(content[0])["text"] == f"scripted context for {marker}", searched.text
requests: Final = upstream_requests(upstream, marker)
assert [value["path"] for value in requests] == [SEARCH_PATH], requests
assert requests[0]["authorization"] == "Bearer integration-provider-key", requests
@pytest.mark.covers("authz.vector_store.list.non_admin_key_access_to_config_store_follows_grants")
def test_non_admin_key_access_to_config_store_follows_grants_after_admin_listing(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
granted: Final = scenario.key(object_permission={"vector_stores": [CONFIG_STORE_ID]})
plain: Final = scenario.key()
assert CONFIG_STORE_ID in listed_ids(gateway)
row: Final = listed_store(gateway, CONFIG_STORE_ID, key=granted)
assert row["is_config"] is True and row["vector_store_name"] == CONFIG_STORE_NAME, row
unlisted: Final = gateway.request("GET", "/vector_store/list", key=plain)
assert unlisted.status_code == 200, unlisted.text
assert CONFIG_STORE_ID not in {value["vector_store_id"] for value in listed_rows(unlisted)}, unlisted.text
for key in (granted, plain):
info = gateway.request("POST", "/vector_store/info", {"vector_store_id": CONFIG_STORE_ID}, key=key)
assert info.status_code == 200, info.text
assert object_value(object_value(info.json())["vector_store"])["is_config"] is True, info.text
forbidden: Final = gateway.request(
"POST", "/vector_store/delete", {"vector_store_id": CONFIG_STORE_ID}, key=granted
)
assert forbidden.status_code in {400, 401, 403}, forbidden.text
assert CONFIG_STORE_ID in listed_ids(gateway)
@pytest.mark.covers("mgmt.vector_store.list.peer_process_keeps_config_store_and_sees_db_store")
def test_peer_process_keeps_config_store_and_sees_db_store_created_elsewhere(gateway: Gateway, peer: Gateway) -> None:
with gateway.scenario() as scenario:
db_store_id: Final = f"vs_db_{uuid.uuid4().hex}"
gateway.post("/vector_store/new", {"vector_store_id": db_store_id, "custom_llm_provider": "openai"})
scenario.cleanups.callback(gateway.request, "POST", "/vector_store/delete", {"vector_store_id": db_store_id})
for side in (gateway, peer, gateway, peer):
assert listed_store(side, CONFIG_STORE_ID)["is_config"] is True
assert listed_store(side, db_store_id)["is_config"] is False
assert config_store_info(side)["is_config"] is True
assert_config_write_refused(side)
gateway.post("/vector_store/delete", {"vector_store_id": db_store_id})
assert db_store_id not in listed_ids(peer)
assert CONFIG_STORE_ID in listed_ids(peer)
@pytest.mark.covers("mgmt.vector_store.chaos.concurrent_burst_keeps_config_store_across_workers")
def test_concurrent_burst_keeps_config_store_and_refuses_every_config_write(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
db_store_ids: Final = tuple(f"vs_db_{uuid.uuid4().hex}" for _ in range(6))
for db_store_id in db_store_ids:
scenario.cleanups.callback(
gateway.request, "POST", "/vector_store/delete", {"vector_store_id": db_store_id}
)
def act(index: int) -> tuple[str, int, str]:
match index % 5:
case 0:
return ("list", *burst_list(gateway))
case 1:
return ("info", *burst_post(gateway, "/vector_store/info", {"vector_store_id": CONFIG_STORE_ID}))
case 2:
return (
"config-update",
*burst_post(
gateway,
"/vector_store/update",
{"vector_store_id": CONFIG_STORE_ID, "vector_store_name": str(index)},
),
)
case 3:
return (
"db-new",
*burst_post(
gateway,
"/vector_store/new",
{
"vector_store_id": db_store_ids[index % len(db_store_ids)],
"custom_llm_provider": "openai",
},
),
)
case _:
return (
"chat",
*burst_post(
gateway,
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"burst {index}"}],
"vector_store_ids": [CONFIG_STORE_ID],
},
),
)
with ThreadPoolExecutor(max_workers=10) as pool:
outcomes: Final = tuple(pool.map(act, range(30)))
expected: Final = {"list": 200, "info": 200, "config-update": 400, "db-new": 200, "chat": 200}
assert [(kind, status) for kind, status, _ in outcomes] == [
(kind, expected[kind]) for kind, _, _ in outcomes
], outcomes
assert all(detail == "config" for kind, _, detail in outcomes if kind == "list"), outcomes
assert listed_store(gateway, CONFIG_STORE_ID)["vector_store_name"] == CONFIG_STORE_NAME
assert config_store_info(gateway)["vector_store_name"] == CONFIG_STORE_NAME
assert store_rows(CONFIG_STORE_ID) == []
assert all(len(store_rows(db_store_id)) == 1 for db_store_id in db_store_ids), "each DB store exactly once"
@pytest.mark.timeout(180)
@pytest.mark.covers("mgmt.vector_store.chaos.redis_outage_keeps_config_store_and_recovers")
def test_redis_outage_keeps_config_store_served_and_recovers(
gateway: Gateway, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
original: Final = os.environ["DATABASE_URL"]
identity: Final = "integration_vs_outage_" + uuid.uuid4().hex
parsed: Final = urlsplit(original)
database_url: Final = urlunsplit((parsed.scheme, parsed.netloc, "/" + identity, "", ""))
with psycopg.connect(original, autocommit=True) as admin:
admin.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(identity)))
try:
with owned_redis(tmp_path) as cache, monkeypatch.context() as environment:
environment.setenv("DATABASE_URL", database_url)
overrides: Final = {
"DATABASE_URL": database_url,
"REDIS_HOST": cache.host,
"REDIS_PORT": str(cache.port),
"REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1",
}
with owned_proxy(gateway, tmp_path, overrides, config=PROXY_CONFIG, workers=2) as candidate:
db_store_id: Final = f"vs_db_{uuid.uuid4().hex}"
for phase in ("before", "during", "after"):
if phase == "during":
cache.stop()
if phase == "after":
cache.start()
for _ in range(4):
assert listed_store(candidate, CONFIG_STORE_ID)["is_config"] is True, phase
assert config_store_info(candidate)["vector_store_name"] == CONFIG_STORE_NAME, phase
assert_config_write_refused(candidate)
created = candidate.request(
"POST",
"/vector_store/new",
{"vector_store_id": f"{db_store_id}_{phase}", "custom_llm_provider": "openai"},
)
assert created.status_code == 200, (phase, created.text)
assert eventually(
lambda phase=phase: store_rows(f"{db_store_id}_{phase}"), lambda rows: len(rows) == 1
), phase
assert f"{db_store_id}_{phase}" in listed_ids(candidate), phase
assert store_rows(CONFIG_STORE_ID) == []
with psycopg.connect(database_url) as fresh:
counted: Final = fresh.execute(
'SELECT count(*) FROM "LiteLLM_ManagedVectorStoresTable" WHERE vector_store_id LIKE %s',
(f"{db_store_id}%",),
).fetchone()
assert counted is not None and counted[0] == 3, counted
finally:
admin.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(identity)))
assert admin.execute("SELECT datname FROM pg_database WHERE datname=%s", (identity,)).fetchall() == []

View file

@ -102,7 +102,7 @@ def test_concurrent_success_and_failure_join_callbacks_and_rows_without_credenti
responses: Final = tuple(pool.map(request, tags))
assert tuple(response.status_code for response in responses) == (200, 400, 200, 400)
assert len(provider.drain()) == 4
batches = []
batches: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls must keep earlier batches
def delivered() -> tuple[dict, ...]:
batches.extend(endpoint.drain())
@ -135,7 +135,7 @@ def test_concurrent_success_and_failure_join_callbacks_and_rows_without_credenti
assert "synthetic callback failure" in json.dumps(event["error_information"])
rows: Final = eventually(
lambda identity=event["id"]: read_rows(
'SELECT request_id, spend, prompt_tokens, completion_tokens, request_tags '
"SELECT request_id, spend, prompt_tokens, completion_tokens, request_tags "
'FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(identity,),
),
@ -156,6 +156,118 @@ def test_concurrent_success_and_failure_join_callbacks_and_rows_without_credenti
assert event["prompt_tokens"] == event["completion_tokens"] == rows[0]["completion_tokens"] == 0
def _responses_frames(identity: str, text: str) -> tuple[bytes, ...]:
output: Final = [
{
"type": "message",
"id": f"msg_{identity}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
]
completed: Final = {
"id": identity,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": output,
"usage": {
"input_tokens": 11,
"output_tokens": 4,
"total_tokens": 15,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens_details": {"reasoning_tokens": 0},
},
}
events: Final = (
{"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}},
{
"type": "response.output_text.delta",
"item_id": f"msg_{identity}",
"output_index": 0,
"content_index": 0,
"delta": text,
},
{"type": "response.completed", "response": completed},
)
return tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events)
@pytest.mark.covers("other.observability.callbacks.streamed_responses_events_carry_provider_response_headers")
def test_streamed_responses_success_callback_carries_provider_apim_request_id(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "resp_" + uuid.uuid4().hex
correlation: Final = "azure-correlation-" + marker
region: Final = "East US 2"
secret: Final = "synthetic-provider-secret-" + marker
sink_secret: Final = "synthetic-sink-secret-" + marker
def upstream(request: Request) -> Reply:
assert request.target.endswith("/responses"), request.target
assert request.headers["authorization"] == f"Bearer {secret}"
assert json.loads(request.body) == {
"model": "gpt-4o-mini",
"input": "header control " + marker,
"stream": True,
}, request.body
return Reply(
content_type="text/event-stream",
chunks=_responses_frames(marker, "streamed control"),
headers={"apim-request-id": correlation, "x-ms-region": region},
)
def sink(request: Request) -> Reply:
assert request.headers["authorization"] == f"Bearer {sink_secret}"
return Reply()
with wire_server(upstream) as provider, wire_server(sink) as endpoint:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update({"callbacks": ["generic_api"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1})
path: Final = tmp_path / "callbacks.yaml"
path.write_text(yaml.safe_dump(config))
with (
owned_proxy(
gateway,
tmp_path,
{
"GENERIC_LOGGER_ENDPOINT": endpoint.url,
"GENERIC_LOGGER_HEADERS": f"Authorization=Bearer {sink_secret}",
},
config=path,
) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key=secret)
response: Final = candidate.request(
"POST", "/v1/responses", {"model": model, "input": "header control " + marker, "stream": True}
)
assert response.status_code == 200, response.text
assert f'"item_id":"msg_{marker}"' in response.text, response.text
assert '"type":"response.completed"' in response.text, response.text
assert len(provider.drain()) == 1
batches: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls must keep earlier batches
def delivered() -> tuple[dict, ...]:
batches.extend(endpoint.drain())
return tuple(
event for batch in batches for event in json.loads(batch.body) if event.get("model_group") == model
)
events: Final = eventually(delivered, lambda values: len(values) == 1, seconds=10)
assert (events[0]["status"], events[0]["stream"], events[0]["call_type"]) == ("success", True, "aresponses")
additional_headers: Final = events[0]["hidden_params"]["additional_headers"] or {}
provider_headers: Final = {
name: value
for name, value in additional_headers.items()
if name in ("llm_provider-apim-request-id", "llm_provider-x-ms-region")
}
assert provider_headers == {
"llm_provider-apim-request-id": correlation,
"llm_provider-x-ms-region": region,
}, json.dumps(events[0]["hidden_params"])
_RAISING_HOOK: Final = """
from litellm.integrations.custom_logger import CustomLogger

View file

@ -5,9 +5,7 @@ from typing import Final
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.client import Gateway
from integration._support.mcp import mcp_peer, register_mcp, tool_names
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
@ -146,6 +144,104 @@ def test_guardrail_denial_prevents_provider_and_preserves_allowed_control(gatewa
assert len(policy.drain()) == 2
@pytest.mark.covers("other.observability.guardrails.bedrock_passthrough_converse_scans_only_caller_content")
def test_bedrock_passthrough_converse_guardrail_ignores_denied_term_in_tool_definition(
gateway: Gateway, tmp_path: Path
) -> None:
identity: Final = "guardrail" + uuid.uuid4().hex
denied: Final = "synthetic denied marker"
allowed: Final = "synthetic allowed weather question"
access_key: Final = "AKIASYNTHETICPASSTHROUGH"
tool_config: Final = {
"tools": [
{
"toolSpec": {
"name": "lookup_weather",
"description": f"Look up the forecast, never answer a {denied}",
"inputSchema": {
"json": {
"type": "object",
"properties": {"city": {"type": "string", "enum": [denied]}},
"required": ["city"],
}
},
}
}
]
}
def guardrail(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api"
texts: Final = json.loads(request.body)["texts"]
result: Final = (
{"action": "BLOCKED", "blocked_reason": "synthetic policy denial"}
if any(denied in text for text in texts)
else {"action": "NONE"}
)
return Reply(body=json.dumps(result).encode())
def runtime(request: Request) -> Reply:
assert request.target == "/model/anthropic.claude-3-haiku-20240307-v1:0/converse"
assert request.headers["authorization"].startswith(f"AWS4-HMAC-SHA256 Credential={access_key}/"), (
request.headers
)
return Reply(
body=json.dumps(
{
"output": {"message": {"role": "assistant", "content": [{"text": "sunny passthrough control"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15},
"metrics": {"latencyMs": 1},
}
).encode()
)
with wire_server(guardrail) as policy, wire_server(runtime) as bedrock, gateway.scenario() as scenario:
model: Final = scenario.model(
model="bedrock/anthropic.claude-3-haiku-20240307-v1:0",
api_key=None,
api_base=bedrock.url,
aws_access_key_id=access_key,
aws_secret_access_key="synthetic-secret",
aws_region_name="us-east-1",
)
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": identity,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "pre_call",
"default_on": True,
"api_base": policy.url,
"api_key": "synthetic-guardrail-key",
},
}
]
path: Final = tmp_path / "bedrock-passthrough.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate:
route: Final = f"/bedrock/model/{model}/converse"
passed: Final = candidate.request(
"POST",
route,
{"messages": [{"role": "user", "content": [{"text": allowed}]}], "toolConfig": tool_config},
)
assert passed.status_code == 200, passed.text
assert passed.json()["output"]["message"]["content"] == [{"text": "sunny passthrough control"}]
forwarded: Final = bedrock.drain()
assert len(forwarded) == 1, "the runtime peer must see exactly the allowed request"
assert json.loads(forwarded[0].body)["toolConfig"] == tool_config
blocked: Final = candidate.request(
"POST",
route,
{"messages": [{"role": "user", "content": [{"text": denied}]}], "toolConfig": tool_config},
)
assert blocked.status_code == 400 and "synthetic policy denial" in blocked.text, blocked.text
assert bedrock.drain() == ()
assert [json.loads(request.body)["texts"] for request in policy.drain()] == [[allowed], [denied]]
@pytest.mark.covers("other.mcp.guardrails.request_selection_blocks_resolved_tool_without_execution")
def test_request_selected_mcp_guardrail_blocks_direct_and_virtual_calls(gateway: Gateway, tmp_path: Path) -> None:
guardrail = "mcp-policy-" + uuid.uuid4().hex

View file

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

View file

@ -1,10 +1,12 @@
from collections.abc import Iterator, Mapping
from typing import Final
from pathlib import Path
import json
import uuid
from collections.abc import Iterator, Mapping
from pathlib import Path
from typing import Final
import pytest
import yaml
from pydantic import JsonValue
from tests.integration._support.client import Gateway, eventually, object_value, string_value
from tests.integration._support.database import read_rows
@ -28,6 +30,31 @@ def test_custom_price_is_reported_and_charged(gateway: Gateway) -> None:
assert params["output_cost_per_token"] == 0.002
@pytest.mark.covers("quota_management.cost_estimate.configured_price.reported_for_model_absent_from_cost_map")
def test_cost_estimate_reports_configured_prices_for_model_absent_from_cost_map(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"openai/integration-on-prem-{uuid.uuid4().hex}",
input_cost_per_token=0.003,
output_cost_per_token=0.007,
)
response: Final = gateway.request(
"POST",
"/cost/estimate",
{"model": model, "input_tokens": 1000, "output_tokens": 500, "num_requests_per_day": 10},
)
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
assert body["input_cost_per_token"] == pytest.approx(0.003), response.text
assert body["output_cost_per_token"] == pytest.approx(0.007), response.text
assert body["input_cost_per_request"] == pytest.approx(1000 * 0.003), response.text
assert body["output_cost_per_request"] == pytest.approx(500 * 0.007), response.text
margin: Final = body["margin_cost_per_request"]
assert isinstance(margin, float), response.text
assert body["cost_per_request"] == pytest.approx(1000 * 0.003 + 500 * 0.007 + margin), response.text
assert body["daily_cost"] == pytest.approx(10 * (1000 * 0.003 + 500 * 0.007 + margin)), response.text
@pytest.mark.covers("quota_management.spend_tracking.default_prices.survive_nullable_sibling_reload")
def test_default_prices_survive_nullable_sibling_and_reload(gateway: Gateway) -> None:
for registration_order in (("custom", "omitted", "nullable"), ("nullable", "omitted", "custom")):
@ -103,6 +130,45 @@ def test_default_prices_survive_nullable_sibling_and_reload(gateway: Gateway) ->
assert float(rows[0]["spend"]) == pytest.approx(expected, rel=1e-6)
COST_MAP_DISPLAY_PRICING_KEYS: Final = frozenset(
{
"input_cost_per_token",
"output_cost_per_token",
"cache_read_input_token_cost",
"cache_creation_input_token_cost",
}
)
def persisted_model_info(identity: str) -> dict[str, JsonValue]:
rows: Final = read_rows('SELECT model_info FROM "LiteLLM_ProxyModelTable" WHERE model_id = %s', (identity,))
assert len(rows) == 1, f"Deployment {identity} has {len(rows)} rows"
stored: Final = rows[0]["model_info"]
return object_value(json.loads(stored) if isinstance(stored, str) else stored)
@pytest.mark.covers("pricing.model_update.echoed_cost_map_price_is_not_persisted_as_override")
def test_saving_echoed_model_info_does_not_freeze_cost_map_price_into_deployment(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
entries: Final = gateway.get("/model/info")["data"]
assert isinstance(entries, list)
target: Final = next(object_value(entry) for entry in entries if object_value(entry)["model_name"] == model)
displayed: Final = object_value(target["model_info"])
identity: Final = string_value(displayed["id"])
assert isinstance(displayed["input_cost_per_token"], float), displayed
assert isinstance(displayed["output_cost_per_token"], float), displayed
fresh: Final = persisted_model_info(identity)
assert {key: value for key, value in fresh.items() if key in COST_MAP_DISPLAY_PRICING_KEYS} == {}, fresh
saved: Final = gateway.request(
"PATCH", f"/model/{identity}/update", {"model_info": {**displayed, "description": "echoed ui save"}}
)
assert saved.status_code == 200, saved.text
stored: Final = persisted_model_info(identity)
assert stored["description"] == "echoed ui save", stored
assert {key: value for key, value in stored.items() if key in COST_MAP_DISPLAY_PRICING_KEYS} == {}, stored
@pytest.mark.covers("quota_management.spend_tracking.default_prices.loaded_router_preserves_cached_defaults")
def test_loaded_router_preserves_cached_defaults_during_real_requests(gateway: Gateway, tmp_path: Path) -> None:
from litellm import Router

View file

@ -0,0 +1,89 @@
import json
import uuid
from typing import Final
import pytest
from tests.integration._support.client import Gateway, eventually, object_value, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.upstream import delete_scenario, register_scenario
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse
INPUT_RATE: Final = 0.001
OUTPUT_RATE: Final = 0.002
CACHE_CREATION_RATE: Final = 0.004
CACHE_READ_RATE: Final = 0.0001
UNCACHED_PROMPT_TOKENS: Final = 1000
CACHE_CREATION_TOKENS: Final = 2000
CACHE_READ_TOKENS: Final = 8000
PROMPT_TOKENS: Final = UNCACHED_PROMPT_TOKENS + CACHE_CREATION_TOKENS + CACHE_READ_TOKENS
COMPLETION_TOKENS: Final = 500
def databricks_cached_response() -> JsonResponse:
return JsonResponse(
content_type="application/json",
body={
"id": "chatcmpl-$REQUEST_ID",
"object": "chat.completion",
"created": 1700000000,
"model": "databricks-claude-integration",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "cached reply"}, "finish_reason": "stop"}
],
"usage": {
"prompt_tokens": PROMPT_TOKENS,
"completion_tokens": COMPLETION_TOKENS,
"total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS,
"cache_creation_input_tokens": CACHE_CREATION_TOKENS,
"cache_read_input_tokens": CACHE_READ_TOKENS,
},
},
)
@pytest.mark.covers("pricing.databricks.cached_prompt_tokens_bill_at_cache_rates")
def test_databricks_cached_prompt_tokens_bill_at_cache_rates_not_input_rate(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"databricks-cache-{uuid.uuid4().hex[:12]}"
handle: Final = register_scenario(scenario_id, databricks_cached_response())
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(
model="databricks/databricks-claude-integration",
api_base=handle.api_base(),
input_cost_per_token=INPUT_RATE,
output_cost_per_token=OUTPUT_RATE,
cache_creation_input_token_cost=CACHE_CREATION_RATE,
cache_read_input_token_cost=CACHE_READ_RATE,
)
response: Final = gateway.request(
"POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "cache control"}]}
)
assert response.status_code == 200, response.text
expected_prompt_cost: Final = (
UNCACHED_PROMPT_TOKENS * INPUT_RATE
+ CACHE_CREATION_TOKENS * CACHE_CREATION_RATE
+ CACHE_READ_TOKENS * CACHE_READ_RATE
)
expected_completion_cost: Final = COMPLETION_TOKENS * OUTPUT_RATE
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(
expected_prompt_cost + expected_completion_cost, rel=1e-6
), response.text
request_id: Final = string_value(object_value(response.json())["id"])
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, metadata, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" '
"WHERE request_id = %s",
(request_id,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["prompt_tokens"] == PROMPT_TOKENS
assert rows[0]["completion_tokens"] == COMPLETION_TOKENS
assert float(rows[0]["spend"]) == pytest.approx(expected_prompt_cost + expected_completion_cost, rel=1e-6)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
breakdown: Final = object_value(parsed["cost_breakdown"])
assert float(breakdown["input_cost"]) == pytest.approx(expected_prompt_cost, rel=1e-6)
assert float(breakdown["output_cost"]) == pytest.approx(expected_completion_cost, rel=1e-6)

View file

@ -0,0 +1,59 @@
import uuid
from typing import Final
import pytest
from tests.integration._support.client import Gateway, eventually, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.upstream import delete_scenario, register_scenario
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse
@pytest.mark.covers("pricing.ocr.annotation_pages_billed_at_annotation_rate")
def test_ocr_annotation_pages_are_billed_at_annotation_cost_per_page(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"ocr-annotation-{uuid.uuid4().hex[:12]}"
handle: Final = register_scenario(
scenario_id,
JsonResponse(
content_type="application/json",
body={
"pages": [{"index": index, "markdown": f"page {index}"} for index in range(3)],
"model": "integration-ocr",
"document_annotation": '{"title": "annotated"}',
"usage_info": {"pages_processed": 3, "pages_processed_annotation": 2, "doc_size_bytes": 4096},
},
),
)
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(
model=f"mistral/integration-ocr-{scenario_id}",
api_base=f"{handle.api_base()}/v1",
ocr_cost_per_page=0.002,
annotation_cost_per_page=0.01,
)
response: Final = gateway.request(
"POST",
"/v1/ocr",
{
"model": model,
"document": {"type": "document_url", "document_url": "https://example.com/annotated.pdf"},
"document_annotation_format": {"type": "json_schema", "json_schema": {"name": "title"}},
},
)
assert response.status_code == 200, response.text
assert response.json()["usage_info"] == {
"pages_processed": 3,
"pages_processed_annotation": 2,
"credits": None,
"doc_size_bytes": 4096,
}, response.text
expected: Final = 3 * 0.002 + 2 * 0.01
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected), response.text
request_id: Final = string_value(response.headers["x-litellm-call-id"])
rows: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (request_id,)),
lambda values: len(values) == 1,
seconds=70,
)
assert float(rows[0]["spend"]) == pytest.approx(expected)

View file

@ -0,0 +1,71 @@
import json
from typing import Final
import httpx
import pytest
from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from tests.integration._support.database import read_rows
STANDARD_INPUT_RATE: Final = 0.001
STANDARD_OUTPUT_RATE: Final = 0.002
ULTRAFAST_INPUT_RATE: Final = 0.01
ULTRAFAST_OUTPUT_RATE: Final = 0.02
def assert_chat_bills_rates(
gateway: Gateway, model: str, service_tier: str | None, input_rate: float, output_rate: float
) -> None:
with httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream:
upstream.get("/__observations").raise_for_status()
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"service tier {service_tier} control"}],
**({} if service_tier is None else {"service_tier": service_tier}),
},
)
assert response.status_code == 200, response.text
expected: Final = 20 * input_rate + 20 * output_rate
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected, rel=1e-6), response.text
observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"]
assert isinstance(observations, list)
assert len(observations) == 1
body: Final = object_value(object_value(observations[0])["body"])
assert body == {
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": f"service tier {service_tier} control"}],
**({} if service_tier is None else {"service_tier": service_tier}),
}, response.text
request_id: Final = string_value(object_value(response.json())["id"])
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, metadata, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id = %s',
(request_id,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["prompt_tokens"] == 20
assert rows[0]["completion_tokens"] == 20
assert float(rows[0]["spend"]) == pytest.approx(expected, rel=1e-6)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
breakdown: Final = object_value(parsed["cost_breakdown"])
assert float(breakdown["input_cost"]) == pytest.approx(20 * input_rate, rel=1e-6)
assert float(breakdown["output_cost"]) == pytest.approx(20 * output_rate, rel=1e-6)
@pytest.mark.covers("quota_management.spend_tracking.service_tier_pricing.ultrafast_bills_ultrafast_rates")
def test_ultrafast_service_tier_bills_ultrafast_rates_and_keeps_pricing_off_the_wire(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(
input_cost_per_token=STANDARD_INPUT_RATE,
output_cost_per_token=STANDARD_OUTPUT_RATE,
input_cost_per_token_ultrafast=ULTRAFAST_INPUT_RATE,
output_cost_per_token_ultrafast=ULTRAFAST_OUTPUT_RATE,
)
assert_chat_bills_rates(gateway, model, "ultrafast", ULTRAFAST_INPUT_RATE, ULTRAFAST_OUTPUT_RATE)
assert_chat_bills_rates(gateway, model, None, STANDARD_INPUT_RATE, STANDARD_OUTPUT_RATE)

View file

@ -0,0 +1,113 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
_ADVISOR_KEY: Final = "synthetic-advisor-key"
_QUESTION: Final = "which index should this query use"
_ADVICE: Final = "use the composite index on (tenant_id, created_at)"
_FINAL_ANSWER: Final = "done, the composite index is the right one"
_ADVISOR_CALL_MESSAGE: Final = {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "advisor-call",
"type": "function",
"function": {"name": "advisor", "arguments": json.dumps({"question": _QUESTION})},
}
],
}
_FINAL_MESSAGE: Final = {"role": "assistant", "content": _FINAL_ANSWER}
def _chat_completion(identity: str, message: dict[str, object], finish_reason: str) -> Reply:
return Reply(
body=json.dumps(
{
"id": f"chatcmpl-{identity}",
"object": "chat.completion",
"created": 1,
"model": "llama-3.3-70b-versatile",
"choices": [{"index": 0, "message": message, "finish_reason": finish_reason}],
"usage": {"prompt_tokens": 10, "completion_tokens": 4, "total_tokens": 14},
}
).encode()
)
def _executor_reply(body: dict[str, object], identity: str) -> Reply:
messages: Final = body["messages"]
assert isinstance(messages, list)
if any(message.get("role") == "tool" for message in messages):
assert messages[-1]["content"] == _ADVICE
return _chat_completion(identity, _FINAL_MESSAGE, "stop")
tools: Final = body["tools"]
assert isinstance(tools, list)
assert tools[0]["function"]["name"] == "advisor"
return _chat_completion(identity, _ADVISOR_CALL_MESSAGE, "tool_calls")
@pytest.mark.covers("providers.anthropic_messages_advisor.sub_call_uses_the_configured_advisor_deployment")
def test_advisor_sub_call_reaches_the_router_deployment_with_its_key_instead_of_anthropic_unauthenticated(
gateway: Gateway,
) -> None:
identity: Final = "advisor-wire-" + uuid.uuid4().hex
def respond(request: Request) -> Reply:
body: Final = json.loads(request.body)
if request.target == "/v1/chat/completions":
assert request.headers["authorization"] == "Bearer integration-provider-key"
return _executor_reply(body, identity)
assert request.target == "/v1/messages"
assert request.headers["x-api-key"] == _ADVISOR_KEY
assert body["model"] == "claude-opus-4-1-20250805"
assert body["messages"] == [
{"role": "user", "content": "please plan the migration"},
{"role": "user", "content": _QUESTION},
]
assert "tools" not in body
return Reply(
body=json.dumps(
{
"id": f"msg-{identity}",
"type": "message",
"role": "assistant",
"model": "claude-opus-4-1-20250805",
"content": [{"type": "text", "text": _ADVICE}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 12, "output_tokens": 6},
}
).encode()
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
executor: Final = scenario.model(model="hosted_vllm/llama-3.3-70b", api_base=wire.url + "/v1")
advisor: Final = scenario.model(
model="anthropic/claude-opus-4-1-20250805", api_base=wire.url, api_key=_ADVISOR_KEY
)
response: Final = gateway.request(
"POST",
"/v1/messages",
{
"model": executor,
"max_tokens": 64,
"messages": [{"role": "user", "content": "please plan the migration"}],
"tools": [{"type": "advisor_20260301", "name": "advisor", "model": advisor}],
},
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert body["content"] == [{"type": "text", "text": _FINAL_ANSWER}], response.text
assert body["stop_reason"] == "end_turn", response.text
assert [request.target for request in wire.drain()] == [
"/v1/chat/completions",
"/v1/messages",
"/v1/chat/completions",
]

View file

@ -0,0 +1,85 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "kimi-k2-thinking"
_API_KEY: Final = "synthetic-azure-ai-key"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_THINKING_BLOCK: Final[JsonValue] = {
"type": "thinking",
"thinking": "The user wants the sum of 17 and 26.",
"signature": "synthetic-signature",
}
_HISTORY_WITH_ANTHROPIC_FIELDS: Final[JsonValue] = [
{
"role": "system",
"content": "You are a calculator.",
"cache_control": {"type": "ephemeral"},
},
{"role": "user", "content": "What is 17 + 26?"},
{
"role": "assistant",
"content": "43",
"thinking_blocks": [_THINKING_BLOCK],
"provider_specific_fields": {"citations": None},
},
{"role": "user", "content": "And doubled?"},
]
_HISTORY_AS_OPENAI_SPEC: Final[JsonValue] = [
{"role": "system", "content": "You are a calculator."},
{"role": "user", "content": "What is 17 + 26?"},
{"role": "assistant", "content": "43"},
{"role": "user", "content": "And doubled?"},
]
def _completion(identity: str) -> bytes:
return json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": _BACKEND,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "86"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 31, "completion_tokens": 2, "total_tokens": 33},
}
).encode()
@pytest.mark.covers("providers.azure_ai.anthropic_message_fields_are_stripped_before_foundry")
def test_azure_ai_strips_thinking_blocks_and_cache_control_from_forwarded_messages(gateway: Gateway) -> None:
identity: Final = f"azure-ai-strip-{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/chat/completions"
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == _BACKEND
assert body["messages"] == _HISTORY_AS_OPENAI_SPEC
return Reply(body=_completion(identity))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"azure_ai/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": _HISTORY_WITH_ANTHROPIC_FIELDS},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["id"] == identity
assert payload["choices"] == [
{
"finish_reason": "stop",
"index": 0,
"message": {"role": "assistant", "content": "86"},
"provider_specific_fields": {},
}
]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]

View file

@ -0,0 +1,48 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_FLEX_MODEL: Final = "azure_ai/FLUX.2-flex"
_PROMPT: Final = "a red fox in the snow"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
@pytest.mark.covers("other.provider_wire.azure_ai.flux2_flex_generation_targets_flex_path_with_bfl_body")
def test_azure_flux2_flex_generation_hits_flex_provider_path_not_pro(gateway: Gateway) -> None:
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/providers/blackforestlabs/v1/flux-2-flex?api-version=preview"
assert request.headers["api-key"] == "synthetic-azure-key"
assert _JSON_OBJECT.validate_json(request.body) == {
"model": "FLUX.2-flex",
"prompt": _PROMPT,
"num_images": 2,
"width": 1536,
"height": 1024,
"guidance": 4.5,
"steps": 32,
}
return Reply(body=json.dumps({"data": [{"b64_json": "aW1n"}, {"b64_json": "aW1n"}]}).encode())
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=_FLEX_MODEL, api_base=wire.url, api_key="synthetic-azure-key", api_version="preview"
)
response: Final = gateway.request(
"POST",
"/v1/images/generations",
{"model": model, "prompt": _PROMPT, "n": 2, "size": "1536x1024", "guidance": 4.5, "steps": 32},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["data"] == [
{"url": None, "b64_json": "aW1n", "revised_prompt": None, "provider_specific_fields": None},
{"url": None, "b64_json": "aW1n", "revised_prompt": None, "provider_specific_fields": None},
]
assert [(request.method, request.target) for request in wire.drain()] == [
("POST", "/providers/blackforestlabs/v1/flux-2-flex?api-version=preview")
]

View file

@ -0,0 +1,46 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
MODEL: Final = "azure_ai/Cohere-rerank-v4.0-fast"
ENTRA_TOKEN: Final = "synthetic-entra-access-token"
QUERY: Final = "which document mentions the gateway"
DOCUMENTS: Final = ("the gateway proxies rerank calls", "unrelated synthetic text")
RESPONSE: Final = json.dumps(
{
"id": "synthetic-rerank-id",
"results": [{"index": 0, "relevance_score": 0.91}, {"index": 1, "relevance_score": 0.03}],
"meta": {"api_version": {"version": "2"}, "billed_units": {"search_units": 1}},
}
).encode()
def entra_rerank_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/providers/cohere/v2/rerank"
assert request.headers["authorization"] == f"Bearer {ENTRA_TOKEN}"
assert "api-key" not in request.headers
body: Final = json.loads(request.body)
assert body == {"model": "Cohere-rerank-v4.0-fast", "query": QUERY, "documents": list(DOCUMENTS), "top_n": 2}
return Reply(body=RESPONSE)
@pytest.mark.covers("other.provider_wire.azure_ai.rerank_entra_token_without_api_key_reaches_provider")
def test_azure_ai_rerank_with_entra_token_and_no_api_key_sends_bearer_to_provider(gateway: Gateway) -> None:
with wire_server(entra_rerank_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=MODEL,
api_key=None,
api_base=f"{wire.url}/providers/cohere/v2",
azure_ad_token=ENTRA_TOKEN,
model_info={"mode": "rerank"},
)
response: Final = gateway.request(
"POST", "/v1/rerank", {"model": model, "query": QUERY, "documents": list(DOCUMENTS), "top_n": 2}
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert [(result["index"], result["relevance_score"]) for result in body["results"]] == [(0, 0.91), (1, 0.03)]
assert len(wire.drain()) == 1, "Expected exactly one provider rerank call"

View file

@ -7,18 +7,22 @@ from typing import Final
import pytest
import yaml
from integration._support.client import Gateway
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
MODEL: Final = "bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0"
TOKEN: Final = "synthetic-bedrock-bearer"
RESPONSE: Final = json.dumps({
"output": {"message": {"role": "assistant", "content": [{"text": "bedrock wire control"}]}},
"stopReason": "end_turn", "usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15},
"metrics": {"latencyMs": 1},
}).encode()
ACCESS_KEY: Final = "AKIAINTEGRATION000002"
CLIENT_OAUTH_TOKEN: Final = "Bearer sk-ant-oat01-synthetic-client-subscription-token"
RESPONSE: Final = json.dumps(
{
"output": {"message": {"role": "assistant", "content": [{"text": "bedrock wire control"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15},
"metrics": {"latencyMs": 1},
}
).encode()
def bearer_peer(request: Request) -> Reply:
@ -34,31 +38,58 @@ def bearer_peer(request: Request) -> Reply:
@pytest.mark.covers("other.provider_wire.bedrock.bearer_sdk_skips_credential_chain")
async def test_bearer_only_sdk_sync_async_requests_do_not_require_aws_credentials(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
async def test_bearer_only_sdk_sync_async_requests_do_not_require_aws_credentials(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
import litellm
empty: Final = tmp_path / "empty-aws-config"
empty.write_text("")
for name in tuple(name for name in os.environ if name.startswith("AWS_")):
monkeypatch.delenv(name, raising=False)
for name, value in {"AWS_CONFIG_FILE": str(empty), "AWS_SHARED_CREDENTIALS_FILE": str(empty), "AWS_EC2_METADATA_DISABLED": "true", "LITELLM_RUST": "false"}.items():
for name, value in {
"AWS_CONFIG_FILE": str(empty),
"AWS_SHARED_CREDENTIALS_FILE": str(empty),
"AWS_EC2_METADATA_DISABLED": "true",
"LITELLM_RUST": "false",
}.items():
monkeypatch.setenv(name, value)
with wire_server(bearer_peer) as wire:
with pytest.raises(litellm.APIConnectionError, match=r"config profile .* could not be found"):
await asyncio.to_thread(litellm.completion, model=MODEL, aws_profile_name="integration-profile-must-not-be-read", aws_region_name="us-east-1", aws_bedrock_runtime_endpoint=wire.url, messages=[{"role": "user", "content": "synthetic credential control"}], timeout=5, num_retries=0)
await asyncio.to_thread(
litellm.completion,
model=MODEL,
aws_profile_name="integration-profile-must-not-be-read",
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint=wire.url,
messages=[{"role": "user", "content": "synthetic credential control"}],
timeout=5,
num_retries=0,
)
assert wire.drain() == ()
for source in ("argument", "environment"):
if source == "environment":
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", TOKEN)
parameters: Final = {
"model": MODEL, "api_key": TOKEN if source == "argument" else None,
"aws_region_name": "us-east-1", "aws_profile_name": "integration-profile-must-not-be-read",
"aws_bedrock_runtime_endpoint": wire.url, "timeout": 5, "num_retries": 0,
"messages": [{"role": "system", "content": "synthetic system"}, {"role": "user", "content": "synthetic bearer request"}],
"model": MODEL,
"api_key": TOKEN if source == "argument" else None,
"aws_region_name": "us-east-1",
"aws_profile_name": "integration-profile-must-not-be-read",
"aws_bedrock_runtime_endpoint": wire.url,
"timeout": 5,
"num_retries": 0,
"messages": [
{"role": "system", "content": "synthetic system"},
{"role": "user", "content": "synthetic bearer request"},
],
"max_tokens": 16,
}
for asynchronous in (False, True):
result: Final = await litellm.acompletion(**parameters) if asynchronous else await asyncio.to_thread(litellm.completion, **parameters)
result: Final = (
await litellm.acompletion(**parameters)
if asynchronous
else await asyncio.to_thread(litellm.completion, **parameters)
)
assert result.choices[0].message.content == "bedrock wire control"
assert result.choices[0].finish_reason == "stop"
assert result.usage.prompt_tokens == 11 and result.usage.completion_tokens == 4
@ -66,28 +97,57 @@ async def test_bearer_only_sdk_sync_async_requests_do_not_require_aws_credential
@pytest.mark.covers("other.provider_wire.bedrock.bearer_db_yaml_survives_reload")
def test_bearer_environment_reference_loads_from_db_and_yaml_and_survives_reload(gateway: Gateway, tmp_path: Path) -> None:
def test_bearer_environment_reference_loads_from_db_and_yaml_and_survives_reload(
gateway: Gateway, tmp_path: Path
) -> None:
empty: Final = tmp_path / "empty-aws-config"
empty.write_text("")
with wire_server(bearer_peer) as wire:
parameters: Final = {
"model": MODEL, "api_key": "os.environ/INTEGRATION_BEARER_TOKEN", "aws_region_name": "us-east-1",
"aws_profile_name": "integration-profile-must-not-be-read", "aws_bedrock_runtime_endpoint": wire.url,
"model": MODEL,
"api_key": "os.environ/INTEGRATION_BEARER_TOKEN",
"aws_region_name": "us-east-1",
"aws_profile_name": "integration-profile-must-not-be-read",
"aws_bedrock_runtime_endpoint": wire.url,
}
alias: Final = f"integration-yaml-{uuid.uuid4().hex}"
configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
configuration["model_list"] = [{"model_name": alias, "litellm_params": parameters, "model_info": {"id": alias}}]
path: Final = tmp_path / "bedrock.yaml"
path.write_text(yaml.safe_dump(configuration))
overrides: Final = {"INTEGRATION_BEARER_TOKEN": TOKEN, "AWS_CONFIG_FILE": str(empty), "AWS_SHARED_CREDENTIALS_FILE": str(empty), "AWS_EC2_METADATA_DISABLED": "true", "LITELLM_RUST": "false"}
with owned_proxy(gateway, tmp_path, overrides, config=path, remove_environment=tuple(name for name in os.environ if name.startswith("AWS_"))) as candidate, candidate.scenario() as scenario:
overrides: Final = {
"INTEGRATION_BEARER_TOKEN": TOKEN,
"AWS_CONFIG_FILE": str(empty),
"AWS_SHARED_CREDENTIALS_FILE": str(empty),
"AWS_EC2_METADATA_DISABLED": "true",
"LITELLM_RUST": "false",
}
with (
owned_proxy(
gateway,
tmp_path,
overrides,
config=path,
remove_environment=tuple(name for name in os.environ if name.startswith("AWS_")),
) as candidate,
candidate.scenario() as scenario,
):
database_model: Final = scenario.model(**parameters)
for generation in range(2):
for model in (alias, database_model):
response: Final = candidate.request("POST", "/v1/chat/completions", {
"model": model, "messages": [{"role": "system", "content": "synthetic system"}, {"role": "user", "content": "synthetic bearer request"}],
"max_tokens": 16, "cache": {"no-cache": True},
})
response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [
{"role": "system", "content": "synthetic system"},
{"role": "user", "content": "synthetic bearer request"},
],
"max_tokens": 16,
"cache": {"no-cache": True},
},
)
assert response.status_code == 200, response.text
assert response.json()["choices"][0]["message"]["content"] == "bedrock wire control"
assert response.json()["usage"]["total_tokens"] == 15
@ -95,5 +155,61 @@ def test_bearer_environment_reference_loads_from_db_and_yaml_and_survives_reload
if generation == 0:
entries: Final = candidate.get("/model/info")["data"]
target: Final = next(entry for entry in entries if entry["model_name"] == database_model)
response: Final = candidate.request("PATCH", f"/model/{target['model_info']['id']}/update", {"model_info": {"description": "bearer reload"}})
response: Final = candidate.request(
"PATCH",
f"/model/{target['model_info']['id']}/update",
{"model_info": {"description": "bearer reload"}},
)
assert response.status_code == 200, response.text
INVOKE_MODEL: Final = "bedrock/invoke/anthropic.claude-3-haiku-20240307-v1:0"
INVOKE_RESPONSE: Final = json.dumps(
{
"id": "msg_synthetic",
"type": "message",
"role": "assistant",
"model": "anthropic.claude-3-haiku-20240307-v1:0",
"content": [{"type": "text", "text": "bedrock invoke wire control"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 11, "output_tokens": 4},
}
).encode()
def sigv4_invoke_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/model/anthropic.claude-3-haiku-20240307-v1:0/invoke"
assert request.headers["authorization"].startswith(f"AWS4-HMAC-SHA256 Credential={ACCESS_KEY}/"), dict(
request.headers
)
assert CLIENT_OAUTH_TOKEN not in request.headers.values(), dict(request.headers)
assert json.loads(request.body)["messages"] == [{"role": "user", "content": "synthetic oauth isolation request"}]
return Reply(body=INVOKE_RESPONSE)
@pytest.mark.covers("providers.bedrock_auth.client_anthropic_oauth_token_never_replaces_sigv4_authorization")
def test_client_anthropic_oauth_authorization_header_does_not_replace_bedrock_sigv4_signature(gateway: Gateway) -> None:
with wire_server(sigv4_invoke_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=INVOKE_MODEL,
api_key=None,
aws_access_key_id=ACCESS_KEY,
aws_secret_access_key="synthetic-secret-key-for-testing",
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint=wire.url,
api_base=wire.url,
)
response: Final = gateway.request(
"POST",
"/v1/messages",
{
"model": model,
"messages": [{"role": "user", "content": "synthetic oauth isolation request"}],
"max_tokens": 16,
},
headers={"Authorization": CLIENT_OAUTH_TOKEN, "x-litellm-api-key": f"Bearer {gateway.key}"},
)
assert response.status_code == 200, response.text
assert response.json()["content"] == [{"type": "text", "text": "bedrock invoke wire control"}], response.text
assert len(wire.drain()) == 1, response.text

View file

@ -0,0 +1,78 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
MODEL: Final = "bedrock/anthropic.claude-3-haiku-20240307-v1:0"
BUCKET: Final = "integration-batch-bucket"
PROMPT: Final = "synthetic completions prompt"
RESPONSES_INPUT: Final = "synthetic responses input"
INPUT_LINES: Final = (
{
"custom_id": "completions-record",
"method": "POST",
"url": "/v1/completions",
"body": {"model": MODEL, "prompt": PROMPT, "max_tokens": 64},
},
{
"custom_id": "responses-record",
"method": "POST",
"url": "/v1/responses",
"body": {"model": MODEL, "input": RESPONSES_INPUT, "max_output_tokens": 16},
},
)
EXPECTED_S3_OBJECT: Final = (
{
"recordId": "completions-record",
"modelInput": {
"messages": [{"role": "user", "content": [{"type": "text", "text": PROMPT}]}],
"max_tokens": 64,
"anthropic_version": "bedrock-2023-05-31",
},
},
{
"recordId": "responses-record",
"modelInput": {
"messages": [{"role": "user", "content": [{"type": "text", "text": RESPONSES_INPUT}]}],
"max_tokens": 16,
"anthropic_version": "bedrock-2023-05-31",
},
},
)
def s3_peer(request: Request) -> Reply:
assert request.method == "PUT" and request.target.startswith(f"/{BUCKET}/"), request.target
assert request.headers["authorization"].startswith("AWS4-HMAC-SHA256 ")
return Reply(body=b"")
@pytest.mark.covers(
"other.provider_wire.bedrock.batch_file_completions_and_responses_records_reach_s3_as_user_messages"
)
def test_completions_and_responses_batch_records_upload_as_anthropic_user_messages(gateway: Gateway) -> None:
with wire_server(s3_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=MODEL,
api_key=None,
api_base=None,
aws_access_key_id="AKIAIOSFODNN7EXAMPLE",
aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
aws_region_name="us-east-1",
s3_bucket_name=BUCKET,
s3_endpoint_url=wire.url,
)
jsonl: Final = "\n".join(json.dumps(line, separators=(",", ":")) for line in INPUT_LINES) + "\n"
response: Final = gateway.request_multipart(
"/v1/files",
{"purpose": "batch", "model": model},
{"file": ("in.jsonl", jsonl.encode(), "application/jsonl")},
)
assert response.status_code == 200, response.text
assert response.json()["object"] == "file" and response.json()["purpose"] == "batch", response.text
uploads: Final = wire.drain()
assert len(uploads) == 1, f"Expected exactly one S3 PUT, saw {[upload.target for upload in uploads]}"
stored: Final = tuple(json.loads(line) for line in uploads[0].body.decode().splitlines() if line.strip())
assert stored == EXPECTED_S3_OBJECT

View file

@ -0,0 +1,60 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
MODEL: Final = "bedrock/invoke/us.anthropic.claude-opus-4-8"
TOKEN: Final = "synthetic-bedrock-bearer"
RESPONSE: Final = json.dumps(
{
"id": "msg_adaptive_control",
"type": "message",
"role": "assistant",
"model": "us.anthropic.claude-opus-4-8",
"content": [{"type": "text", "text": "adaptive thinking control"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 12, "output_tokens": 5},
}
).encode()
def adaptive_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/model/us.anthropic.claude-opus-4-8/invoke"
assert request.headers["authorization"] == f"Bearer {TOKEN}"
body: Final = json.loads(request.body)
assert body["messages"] == [{"role": "user", "content": [{"type": "text", "text": "synthetic effort request"}]}]
assert body["thinking"]["type"] == "adaptive", body
assert body["output_config"] == {"effort": "high"}, body
assert "budget_tokens" not in json.dumps(body), body
return Reply(body=RESPONSE)
@pytest.mark.covers("other.provider_wire.bedrock.prefixed_opus_4_8_reasoning_effort_sends_adaptive_thinking")
def test_prefixed_opus_4_8_reasoning_effort_reaches_bedrock_as_adaptive_thinking_not_budget_tokens(
gateway: Gateway,
) -> None:
with wire_server(adaptive_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=MODEL,
api_key=TOKEN,
aws_region_name="us-east-1",
api_base=wire.url,
aws_bedrock_runtime_endpoint=wire.url,
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": "synthetic effort request"}],
"max_tokens": 4096,
"reasoning_effort": "high",
},
)
assert response.status_code == 200, response.text
assert response.json()["choices"][0]["message"]["content"] == "adaptive thinking control"
assert response.json()["usage"]["prompt_tokens"] == 12 and response.json()["usage"]["completion_tokens"] == 5
assert len(wire.drain()) == 1

View file

@ -0,0 +1,46 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from integration.providers.test_bedrock_auth_wire import MODEL, RESPONSE, TOKEN
GUARDRAIL: Final = {"guardrailIdentifier": "integration-guardrail", "guardrailVersion": "DRAFT", "trace": "enabled"}
PERFORMANCE: Final = {"latency": "optimized"}
def converse_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/model/anthropic.claude-3-haiku-20240307-v1%3A0/converse"
body: Final = json.loads(request.body)
assert body["inferenceConfig"] == {"maxTokens": 16, "temperature": 0.2}, body
assert body["guardrailConfig"] == GUARDRAIL, body
assert body["performanceConfig"] == PERFORMANCE, body
return Reply(body=RESPONSE)
@pytest.mark.covers("other.provider_wire.bedrock.converse_config_blocks_sent_once_at_top_level")
def test_guardrail_and_performance_config_are_not_duplicated_inside_inference_config(gateway: Gateway) -> None:
with wire_server(converse_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=MODEL,
api_key=TOKEN,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint=wire.url,
guardrailConfig=GUARDRAIL,
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": "synthetic guardrail request"}],
"max_tokens": 16,
"temperature": 0.2,
"performanceConfig": PERFORMANCE,
"cache": {"no-cache": True},
},
)
assert response.status_code == 200, response.text
assert response.json()["choices"][0]["message"]["content"] == "bedrock wire control"
assert len(wire.drain()) == 1, response.text

View file

@ -0,0 +1,58 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
MODEL: Final = "bedrock/cohere.embed-english-v3"
TOKEN: Final = "synthetic-bedrock-bearer"
INPUT: Final = "hello world"
VECTOR: Final = [0.1, 0.2, 0.3]
RESPONSE: Final = json.dumps(
{
"embeddings": {"float": [VECTOR]},
"id": "synthetic-cohere-embed",
"response_type": "embeddings_by_type",
"texts": [INPUT],
}
).encode()
def cohere_english_v3_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/model/cohere.embed-english-v3/invoke"
assert request.headers["authorization"] == f"Bearer {TOKEN}"
assert json.loads(request.body) == {
"texts": [INPUT],
"input_type": "search_document",
"embedding_types": ["float"],
"output_dimension": 512,
}
return Reply(body=RESPONSE)
@pytest.mark.covers("other.provider_wire.bedrock.cohere_embed_english_v3_accepts_encoding_format")
def test_cohere_embed_english_v3_accepts_encoding_format_and_dimensions(gateway: Gateway) -> None:
with wire_server(cohere_english_v3_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=MODEL,
api_key=TOKEN,
api_base=wire.url,
aws_region_name="us-east-1",
)
for encoding_format in ("float", "base64"):
response: Final = gateway.request(
"POST",
"/v1/embeddings",
{
"model": model,
"input": INPUT,
"encoding_format": encoding_format,
"dimensions": 512,
},
)
assert response.status_code == 200, f"encoding_format={encoding_format}: {response.text}"
assert response.json()["data"] == [
{"object": "embedding", "index": 0, "embedding": VECTOR, "type": "float"},
], response.text
assert len(wire.drain()) == 1, f"encoding_format={encoding_format} never reached Bedrock"

View file

@ -0,0 +1,50 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
MODEL: Final = "bedrock/converse/us.openai.gpt-5.6-sol"
TOKEN: Final = "synthetic-bedrock-bearer"
RESPONSE: Final = json.dumps(
{
"output": {"message": {"role": "assistant", "content": [{"text": "gpt-5 reasoning wire control"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 9, "outputTokens": 5, "totalTokens": 14},
"metrics": {"latencyMs": 1},
}
).encode()
def gpt5_converse_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/model/us.openai.gpt-5.6-sol/converse"
assert request.headers["authorization"] == f"Bearer {TOKEN}"
body: Final = json.loads(request.body)
assert body["messages"] == [{"role": "user", "content": [{"text": "synthetic reasoning request"}]}]
assert body["additionalModelRequestFields"] == {"reasoning": {"effort": "high"}}, body
assert body["inferenceConfig"] == {"maxTokens": 16}, body
return Reply(body=RESPONSE)
@pytest.mark.covers("providers.bedrock_converse.gpt5_reasoning_effort_reaches_provider_as_reasoning_effort")
def test_gpt5_reasoning_effort_is_accepted_and_sent_as_converse_reasoning_effort(gateway: Gateway) -> None:
with wire_server(gpt5_converse_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=MODEL, api_key=TOKEN, aws_region_name="us-east-1", aws_bedrock_runtime_endpoint=wire.url
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": "synthetic reasoning request"}],
"reasoning_effort": "high",
"max_tokens": 16,
"cache": {"no-cache": True},
},
)
assert response.status_code == 200, response.text
assert response.json()["choices"][0]["message"]["content"] == "gpt-5 reasoning wire control", response.text
assert response.json()["usage"]["total_tokens"] == 14, response.text
assert len(wire.drain()) == 1

View file

@ -0,0 +1,88 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
MODEL_ID: Final = "us.anthropic.claude-sonnet-5"
TOKEN: Final = "synthetic-bedrock-bearer"
TOOL_SEARCH_TOOL: Final = {"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"}
DEFERRED_TOOL: Final = {
"name": "get_weather",
"description": "Weather lookup",
"input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
"defer_loading": True,
}
RESPONSE: Final = json.dumps(
{
"id": "msg_tool_search_control",
"type": "message",
"role": "assistant",
"model": MODEL_ID,
"content": [
{
"type": "server_tool_use",
"id": "srvtoolu_control",
"name": "tool_search_tool_regex",
"input": {"pattern": "weather"},
},
{
"type": "tool_search_tool_result",
"tool_use_id": "srvtoolu_control",
"content": {
"type": "tool_search_tool_search_result",
"tool_references": [{"type": "tool_reference", "tool_name": "get_weather"}],
},
},
{"type": "text", "text": "tool search wire control"},
],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 12, "output_tokens": 6},
}
).encode()
def tool_search_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == f"/model/{MODEL_ID}/invoke", request.target
assert request.headers["authorization"] == f"Bearer {TOKEN}"
body: Final = json.loads(request.body)
assert body["anthropic_beta"] == ["tool-search-tool-2025-10-19"], body
assert body["messages"] == [{"role": "user", "content": "find the weather tool"}]
assert body["tools"] == [TOOL_SEARCH_TOOL, DEFERRED_TOOL], body["tools"]
assert body["max_tokens"] == 64
assert "model" not in body
return Reply(body=RESPONSE)
@pytest.mark.covers("providers.bedrock_invoke.tool_search_gen5_claude_sends_bedrock_beta_and_reports_support")
def test_gen5_claude_bedrock_invoke_messages_tool_search_sends_bedrock_beta_field(gateway: Gateway) -> None:
with wire_server(tool_search_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"bedrock/invoke/{MODEL_ID}",
api_key=TOKEN,
aws_region_name="us-east-1",
api_base=wire.url,
)
response: Final = gateway.request(
"POST",
"/v1/messages",
{
"model": model,
"max_tokens": 64,
"messages": [{"role": "user", "content": "find the weather tool"}],
"tools": [TOOL_SEARCH_TOOL, DEFERRED_TOOL],
},
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert body["content"][2] == {"type": "text", "text": "tool search wire control"}, response.text
assert body["stop_reason"] == "end_turn"
assert body["usage"]["input_tokens"] == 12 and body["usage"]["output_tokens"] == 6
assert len(wire.drain()) == 1
entries: Final = gateway.get("/v1/model/info")["data"]
assert isinstance(entries, list)
info: Final = next(entry for entry in entries if isinstance(entry, dict) and entry["model_name"] == model)
assert isinstance(info["model_info"], dict)
assert info["model_info"]["supports_tool_search"] is True, info["model_info"]

View file

@ -0,0 +1,88 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
MODEL: Final = "bedrock_mantle/openai.gpt-5.6-sol"
TOKEN: Final = "synthetic-mantle-bearer"
CIPHERTEXT: Final = "synthetic-compaction-ciphertext"
CALL_ID: Final = "call_synthetic_shell"
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
ACTION: Final[dict[str, JsonValue]] = {"type": "exec", "command": ["ls", "-la"], "timeout_ms": 1000}
RESPONSE: Final = json.dumps(
{
"id": "resp_synthetic_mantle",
"object": "response",
"created_at": 1789788253,
"status": "completed",
"model": "openai.gpt-5.6-sol",
"output": [
{
"type": "message",
"id": "msg_synthetic_mantle",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "mantle wire control", "annotations": []}],
}
],
"usage": {"input_tokens": 21, "output_tokens": 4, "total_tokens": 25},
}
).encode()
def user_turn(text: str) -> JsonValue:
return {"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]}
def codex_history(marker: str) -> tuple[JsonValue, ...]:
return (
user_turn(f"first turn {marker}"),
{"type": "agent_message", "role": "assistant", "content": [{"type": "output_text", "text": "sub-agent reply"}]},
{"type": "context_compaction", "encrypted_content": CIPHERTEXT},
{"type": "local_shell_call", "call_id": CALL_ID, "status": "completed", "action": ACTION},
{"type": "function_call_output", "call_id": CALL_ID, "output": "synthetic shell output"},
user_turn(f"next turn {marker}"),
)
def mantle_history(marker: str) -> tuple[JsonValue, ...]:
return (
user_turn(f"first turn {marker}"),
{"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "sub-agent reply"}]},
{"type": "compaction", "encrypted_content": CIPHERTEXT},
{"type": "function_call", "call_id": CALL_ID, "name": "local_shell", "arguments": json.dumps(ACTION)},
{"type": "function_call_output", "call_id": CALL_ID, "output": "synthetic shell output"},
user_turn(f"next turn {marker}"),
)
@pytest.mark.covers("other.provider_wire.bedrock_mantle.codex_history_items_reach_mantle_as_supported_types")
def test_codex_agent_message_context_compaction_and_local_shell_call_reach_mantle_as_supported_items(
gateway: Gateway,
) -> None:
marker: Final = uuid.uuid4().hex
expected_input: Final = list(mantle_history(marker))
def mantle_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/openai/v1/responses", request.target
assert request.headers["authorization"] == f"Bearer {TOKEN}"
body: Final = JSON_OBJECT.validate_json(request.body)
assert body["model"] == "openai.gpt-5.6-sol", body
assert body["input"] == expected_input, body["input"]
return Reply(body=RESPONSE)
with wire_server(mantle_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=MODEL, api_key=TOKEN, api_base=wire.url, aws_region_name="us-east-2")
response: Final = gateway.request(
"POST", "/v1/responses", {"model": model, "input": list(codex_history(marker)), "store": False}
)
assert response.status_code == 200, response.text
assert response.json()["output"][0]["content"][0]["text"] == "mantle wire control", response.text
assert response.json()["usage"]["total_tokens"] == 25, response.text
forwarded: Final = wire.drain()
assert len(forwarded) == 1, forwarded
assert JSON_OBJECT.validate_json(forwarded[0].body)["input"] == expected_input, forwarded[0].body

View file

@ -0,0 +1,106 @@
import json
from collections.abc import Callable
from typing import Final
from uuid import uuid4
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_MODEL: Final = "bedrock_mantle/openai.gpt-5.6-sol"
_TOKEN: Final = "synthetic-mantle-bearer"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_SHELL_ACTION: Final[dict[str, JsonValue]] = {"type": "exec", "command": ["ls", "-la"], "timeout_ms": 1000}
_OUTPUT_MESSAGE: Final[dict[str, JsonValue]] = {
"type": "message",
"id": "msg_mantle",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "mantle wire control", "annotations": []}],
}
_RESPONSE: Final = json.dumps(
{
"id": "resp_mantle",
"object": "response",
"status": "completed",
"created_at": 1700000000,
"model": "gpt-5.6-sol",
"output": [_OUTPUT_MESSAGE],
"usage": {
"input_tokens": 11,
"output_tokens": 4,
"total_tokens": 15,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens_details": {"reasoning_tokens": 0},
},
}
).encode()
def _codex_history(marker: str) -> list[JsonValue]:
return [
{"type": "message", "role": "user", "content": f"delegate to a subagent {marker}"},
{
"type": "agent_message",
"id": "msg_agent",
"content": [{"type": "text", "text": "sub-agent said "}, {"type": "text", "encrypted_content": "hello"}],
},
{"type": "context_compaction", "id": "cmp_1", "encrypted_content": "compacted-history"},
{
"type": "local_shell_call",
"id": "lsc_1",
"call_id": "call_shell",
"status": "completed",
"action": _SHELL_ACTION,
},
{"type": "function_call_output", "call_id": "call_shell", "output": "total 0"},
]
def _mantle_history(marker: str) -> list[JsonValue]:
return [
{"type": "message", "role": "user", "content": f"delegate to a subagent {marker}"},
{"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": "sub-agent said hello"}]},
{"type": "compaction", "encrypted_content": "compacted-history"},
{
"type": "function_call",
"call_id": "call_shell",
"name": "local_shell",
"arguments": json.dumps(_SHELL_ACTION),
},
{"type": "function_call_output", "call_id": "call_shell", "output": "total 0"},
]
def _mantle_peer(marker: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/openai/v1/responses", request.target
assert request.headers["authorization"] == f"Bearer {_TOKEN}"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["input"] == _mantle_history(marker), json.dumps(body["input"])
return Reply(body=_RESPONSE)
return respond
@pytest.mark.covers("providers.bedrock_mantle.codex_history_items_reach_the_wire_as_supported_input_items")
def test_codex_agent_message_compaction_and_local_shell_items_are_rewritten_for_mantle(gateway: Gateway) -> None:
marker: Final = uuid4().hex
with wire_server(_mantle_peer(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_TOKEN, aws_region_name="us-east-1")
response: Final = gateway.request(
"POST", "/v1/responses", {"model": model, "input": _codex_history(marker), "stream": False}
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["output"] == [
{
**_OUTPUT_MESSAGE,
"phase": None,
"content": [
{"type": "output_text", "text": "mantle wire control", "annotations": [], "logprobs": None}
],
}
], response.text
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/openai/v1/responses")]

View file

@ -0,0 +1,53 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "openai.gpt-5.6-sol"
_API_KEY: Final = "synthetic-mantle-bearer"
_PROMPT: Final = "synthetic long conversation control"
_PROMPT_TOKENS: Final = 1055489
_MODEL_MAXIMUM: Final = 1050000
_RESPONSES_PATH: Final = "/openai/v1/responses"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_OVERFLOW_BODY: Final = json.dumps(
{
"error": {
"code": "validation_error",
"message": f"prompt tokens ({_PROMPT_TOKENS}) exceed model maximum ({_MODEL_MAXIMUM}) for {_BACKEND}",
"type": "invalid_request_error",
}
}
).encode()
def _overflow_peer(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == _RESPONSES_PATH
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == _BACKEND
assert _PROMPT in json.dumps(body["input"]), body
return Reply(status=400, body=_OVERFLOW_BODY)
@pytest.mark.covers("other.provider_wire.bedrock_mantle.context_overflow_is_reported_as_prompt_too_long")
def test_bedrock_mantle_context_overflow_returns_400_saying_prompt_is_too_long(gateway: Gateway) -> None:
with wire_server(_overflow_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"bedrock_mantle/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}]},
)
assert response.status_code == 400, response.text
error: Final = _JSON_OBJECT.validate_json(response.content)["error"]
assert isinstance(error, dict), response.text
assert error["code"] == "400", response.text
message: Final = error["message"]
assert isinstance(message, str), response.text
assert f"prompt is too long: {_PROMPT_TOKENS} tokens > {_MODEL_MAXIMUM} maximum" in message, response.text
assert [(request.method, request.target) for request in wire.drain()] == [("POST", _RESPONSES_PATH)]

View file

@ -0,0 +1,88 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
BEDROCK_MODEL: Final = "us.anthropic.claude-opus-5-v1:0"
TOKEN: Final = "synthetic-bedrock-bearer"
SNIPPET: Final = "synthetic snippet about the integration harness"
INTERCEPTED_TURN: Final = (
{"type": "server_tool_use", "id": "srvtoolu_synthetic", "name": "web_search", "input": {"query": "harness docs"}},
{
"type": "web_search_tool_result",
"tool_use_id": "srvtoolu_synthetic",
"content": [
{
"type": "web_search_result",
"url": "https://example.test/harness",
"title": "Harness",
"page_age": None,
"encrypted_content": "",
"snippet": SNIPPET,
},
],
},
{"type": "text", "text": "The harness is documented at example.test"},
)
FLATTENED_TURN: Final = (
{
"type": "text",
"text": f"Web search results for 'harness docs':\n\nTitle: Harness\nURL: https://example.test/harness\nSnippet: {SNIPPET}",
},
{"type": "text", "text": "The harness is documented at example.test"},
)
REPLY: Final = json.dumps(
{
"id": "msg_synthetic_replay",
"type": "message",
"role": "assistant",
"model": BEDROCK_MODEL,
"content": [{"type": "text", "text": "replay accepted"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 30, "output_tokens": 3},
}
).encode()
def bedrock_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == f"/model/{BEDROCK_MODEL}/invoke"
assert request.headers["authorization"] == f"Bearer {TOKEN}"
body: Final = json.loads(request.body)
assert body["messages"] == [
{"role": "user", "content": "where is the harness documented"},
{"role": "assistant", "content": list(FLATTENED_TURN)},
{"role": "user", "content": "and what does it say"},
], request.body.decode()
assert "tools" not in body, request.body.decode()
return Reply(body=REPLY)
@pytest.mark.covers("providers.bedrock_messages.replayed_intercepted_web_search_turn_is_flattened_to_text")
def test_replayed_intercepted_web_search_turn_reaches_bedrock_as_text_and_answers(gateway: Gateway) -> None:
with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"bedrock/{BEDROCK_MODEL}",
api_key=TOKEN,
api_base=wire.url,
aws_region_name="us-east-1",
)
response: Final = gateway.request(
"POST",
"/v1/messages",
{
"model": model,
"max_tokens": 64,
"messages": [
{"role": "user", "content": "where is the harness documented"},
{"role": "assistant", "content": list(INTERCEPTED_TURN)},
{"role": "user", "content": "and what does it say"},
],
},
headers={"x-api-key": gateway.key, "anthropic-version": "2023-06-01"},
)
assert response.status_code == 200, response.text
assert response.json()["content"] == [{"type": "text", "text": "replay accepted"}], response.text
assert len(wire.drain()) == 1

View file

@ -0,0 +1,42 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.upstream import _aws_event_frame
from integration._support.wire import Reply, Request, wire_server
_MODEL_ID: Final = "anthropic.claude-sonnet-5-v1:0"
_EVENT_STREAM: Final = "application/vnd.amazon.eventstream"
_REQUEST_BODY: Final = {"messages": [{"role": "user", "content": [{"text": "synthetic passthrough stream"}]}]}
_EVENTS: Final = (
("messageStart", {"role": "assistant"}),
("contentBlockDelta", {"delta": {"text": "bedrock stream control"}, "contentBlockIndex": 0}),
("messageStop", {"stopReason": "end_turn"}),
("metadata", {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}),
)
_STREAM_BYTES: Final = b"".join(_aws_event_frame(kind, payload, "sc", "u") for kind, payload in _EVENTS)
def event_stream_peer(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == f"/model/{_MODEL_ID}/converse-stream"
assert json.loads(request.body)["messages"] == _REQUEST_BODY["messages"]
return Reply(body=_STREAM_BYTES, content_type=_EVENT_STREAM)
@pytest.mark.covers("other.provider_wire.bedrock.passthrough_stream_keeps_event_stream_content_type")
def test_bedrock_passthrough_converse_stream_response_carries_event_stream_content_type(gateway: Gateway) -> None:
with wire_server(event_stream_peer) as wire, gateway.scenario() as scenario:
deployment: Final = scenario.model(
model=f"bedrock/{_MODEL_ID}",
api_base=wire.url,
aws_access_key_id="AKIASCRIPTEDPROVIDER",
aws_secret_access_key="scripted-secret",
aws_region_name="us-east-1",
)
response: Final = gateway.request("POST", f"/bedrock/model/{deployment}/converse-stream", _REQUEST_BODY)
assert response.status_code == 200, response.text
assert len(wire.drain()) == 1, response.text
assert response.headers.get("content-type") == _EVENT_STREAM, dict(response.headers)
assert response.content == _STREAM_BYTES, response.text

View file

@ -0,0 +1,92 @@
import json
import os
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
MODEL: Final = "bedrock/arn:aws:bedrock:us-east-1::foundation-model/cohere.rerank-v3-5:0"
ACCESS_KEY: Final = "AKIAINTEGRATION000002"
FORWARDED_FOR: Final = "203.0.113.5"
RESPONSE: Final = json.dumps(
{"results": [{"index": 1, "relevanceScore": 0.9}, {"index": 0, "relevanceScore": 0.1}]}
).encode()
def signed_headers(authorization: str) -> tuple[str, ...]:
return tuple(authorization.split("SignedHeaders=")[1].split(",")[0].split(";"))
def rerank_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/rerank"
assert request.headers["authorization"].startswith(f"AWS4-HMAC-SHA256 Credential={ACCESS_KEY}/")
assert signed_headers(request.headers["authorization"]) == ("content-type", "host", "x-amz-date"), request.headers[
"authorization"
]
assert request.headers["x-forwarded-for"] == FORWARDED_FOR
body: Final = json.loads(request.body)
assert body["queries"] == [{"textQuery": {"text": "synthetic rerank query"}, "type": "TEXT"}]
assert body["rerankingConfiguration"]["bedrockRerankingConfiguration"]["modelConfiguration"] == {
"modelArn": "arn:aws:bedrock:us-east-1::foundation-model/cohere.rerank-v3-5:0"
}
assert body["rerankingConfiguration"]["bedrockRerankingConfiguration"]["numberOfResults"] == 2
return Reply(body=RESPONSE)
@pytest.mark.covers("providers.bedrock_rerank.forwarded_client_headers_are_sent_unsigned")
def test_forwarded_client_header_on_rerank_is_excluded_from_the_sigv4_signature(
gateway: Gateway, tmp_path: Path
) -> None:
empty: Final = tmp_path / "empty-aws-config"
empty.write_text("")
configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
configuration["general_settings"]["forward_client_headers_to_llm_api"] = True
path: Final = tmp_path / "forwarding.yaml"
path.write_text(yaml.safe_dump(configuration))
overrides: Final = {
"AWS_CONFIG_FILE": str(empty),
"AWS_SHARED_CREDENTIALS_FILE": str(empty),
"AWS_EC2_METADATA_DISABLED": "true",
"LITELLM_RUST": "false",
}
with wire_server(rerank_peer) as wire:
with (
owned_proxy(
gateway,
tmp_path,
overrides,
config=path,
remove_environment=tuple(name for name in os.environ if name.startswith("AWS_")),
) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(
model=MODEL,
api_key=None,
api_base=None,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint=wire.url,
aws_access_key_id=ACCESS_KEY,
aws_secret_access_key="synthetic-rerank-secret-key-for-testing",
)
response: Final = candidate.request(
"POST",
"/v1/rerank",
{
"model": model,
"query": "synthetic rerank query",
"documents": ["first synthetic document", "second synthetic document"],
"top_n": 2,
},
headers={"x-forwarded-for": FORWARDED_FOR},
)
assert response.status_code == 200, response.text
assert response.json()["results"] == [
{"index": 1, "relevance_score": 0.9},
{"index": 0, "relevance_score": 0.1},
], response.text
assert len(wire.drain()) == 1

View file

@ -0,0 +1,94 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
MODEL: Final = "bedrock/converse/global.anthropic.claude-opus-4-8"
TOKEN: Final = "synthetic-bedrock-bearer"
PROMPT: Final = "How many prime numbers are less than 30? Think it through, then answer with just the number."
RESPONSES_PROMPT: Final = "How many prime numbers are less than 30? Answer with just the number."
REDACTED_DATA: Final = "RWRhY3RlZC1ieS1CZWRyb2Nr"
INPUT_TOKENS: Final = 31
OUTPUT_TOKENS: Final = 257
RESPONSE: Final = json.dumps(
{
"output": {
"message": {
"role": "assistant",
"content": [{"reasoningContent": {"redactedContent": REDACTED_DATA}}, {"text": "10"}],
}
},
"stopReason": "end_turn",
"usage": {
"inputTokens": INPUT_TOKENS,
"outputTokens": OUTPUT_TOKENS,
"totalTokens": INPUT_TOKENS + OUTPUT_TOKENS,
},
"metrics": {"latencyMs": 1},
}
).encode()
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_JSON_LIST: Final = TypeAdapter(list[dict[str, JsonValue]])
def redacted_thinking_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/model/global.anthropic.claude-opus-4-8/converse"
assert request.headers["authorization"] == f"Bearer {TOKEN}"
body: Final = json.loads(request.body)
assert body["messages"] in (
[{"role": "user", "content": [{"text": PROMPT}]}],
[{"role": "user", "content": [{"text": RESPONSES_PROMPT}]}],
), body
assert body["additionalModelRequestFields"]["thinking"]["type"] == "adaptive", body
return Reply(body=RESPONSE)
@pytest.mark.covers("other.provider_wire.bedrock.hidden_thinking_tokens_are_not_reported_as_text")
def test_bedrock_redacted_thinking_is_not_reported_as_zero_reasoning_tokens(gateway: Gateway) -> None:
with wire_server(redacted_thinking_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=MODEL, api_key=TOKEN, aws_region_name="us-east-1", aws_bedrock_runtime_endpoint=wire.url
)
chat: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": PROMPT}],
"max_tokens": 4000,
"reasoning_effort": "max",
},
)
assert chat.status_code == 200, chat.text
chat_body: Final = _JSON_OBJECT.validate_json(chat.content)
message: Final = _JSON_OBJECT.validate_python(_JSON_LIST.validate_python(chat_body["choices"])[0]["message"])
assert message["content"] == "10", chat.text
assert message["thinking_blocks"] == [{"type": "redacted_thinking", "data": REDACTED_DATA}], chat.text
usage: Final = _JSON_OBJECT.validate_python(chat_body["usage"])
assert usage["completion_tokens"] == OUTPUT_TOKENS, chat.text
details: Final = _JSON_OBJECT.validate_python(usage["completion_tokens_details"])
assert details == {}, chat.text
assert len(wire.drain()) == 1
responses: Final = gateway.request(
"POST",
"/v1/responses",
{"model": model, "input": RESPONSES_PROMPT, "max_output_tokens": 4000, "reasoning": {"effort": "max"}},
)
assert responses.status_code == 200, responses.text
responses_body: Final = _JSON_OBJECT.validate_json(responses.content)
output: Final = _JSON_LIST.validate_python(responses_body["output"])
reasoning_items: Final = tuple(item for item in output if item["type"] == "reasoning")
assert len(reasoning_items) == 1, responses.text
assert reasoning_items[0]["encrypted_content"] == json.dumps(
[{"type": "redacted_thinking", "data": REDACTED_DATA}], separators=(",", ":")
), responses.text
responses_usage: Final = _JSON_OBJECT.validate_python(responses_body["usage"])
assert responses_usage["output_tokens"] == OUTPUT_TOKENS, responses.text
assert _JSON_OBJECT.validate_python(responses_usage["output_tokens_details"])["reasoning_tokens"] == 0, (
responses.text
)
assert len(wire.drain()) == 1

View file

@ -0,0 +1,62 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "qwen3.7-plus"
_API_KEY: Final = "synthetic-dashscope-key"
_PROMPT: Final = "What is 3^3?"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
def _completion(identity: str) -> bytes:
return json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": _BACKEND,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "27"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 17, "completion_tokens": 5, "total_tokens": 22},
}
).encode()
@pytest.mark.covers("other.provider_wire.dashscope.reasoning_effort_reaches_provider")
def test_dashscope_chat_forwards_reasoning_effort_none_to_the_provider(gateway: Gateway) -> None:
identity: Final = f"dashscope-reasoning-{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/chat/completions"
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
assert _JSON_OBJECT.validate_json(request.body) == {
"model": _BACKEND,
"messages": [{"role": "user", "content": _PROMPT}],
"reasoning_effort": "none",
}
return Reply(body=_completion(identity))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"dashscope/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}], "reasoning_effort": "none"},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["id"] == identity
assert payload["choices"] == [
{
"finish_reason": "stop",
"index": 0,
"message": {"role": "assistant", "content": "27", "provider_specific_fields": {"refusal": None}},
"provider_specific_fields": {},
}
]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]

View file

@ -0,0 +1,129 @@
import json
import uuid
from collections.abc import Mapping
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
_BACKEND: Final = "databricks-glm-5-2"
_API_KEY: Final = "synthetic-databricks-key"
_PROMPT: Final = "Summarise the cached briefing in one sentence."
_PROVIDER_USAGE: Final[Mapping[str, JsonValue]] = {
"prompt_tokens": 12011,
"completion_tokens": 8,
"total_tokens": 12019,
"cache_read_input_tokens": 12002,
"cache_creation_input_tokens": 0,
}
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
class _PromptTokensDetails(BaseModel):
model_config = ConfigDict(extra="ignore")
cached_tokens: int | None = None
class _Usage(BaseModel):
model_config = ConfigDict(extra="ignore")
prompt_tokens: int
completion_tokens: int
total_tokens: int
prompt_tokens_details: _PromptTokensDetails | None = None
class _Delta(BaseModel):
model_config = ConfigDict(extra="ignore")
content: str | None = None
class _Choice(BaseModel):
model_config = ConfigDict(extra="ignore")
delta: _Delta
class _Chunk(BaseModel):
model_config = ConfigDict(extra="ignore")
id: str
choices: tuple[_Choice, ...]
usage: _Usage | None = None
def _frame(identity: str, choices: list[Mapping[str, object]], usage: Mapping[str, JsonValue] | None = None) -> bytes:
value: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": _BACKEND,
"choices": choices,
**({} if usage is None else {"usage": usage}),
}
return b"data: " + json.dumps(value).encode() + b"\n\n"
@pytest.mark.covers("other.provider_wire.databricks.stream_usage_and_cache_reads_reach_client_and_spend_log")
def test_databricks_stream_final_usage_chunk_reaches_client_and_spend_log(gateway: Gateway) -> None:
identity: Final = f"databricks-stream-{uuid.uuid4().hex}"
frames: Final = (
_frame(
identity, [{"index": 0, "delta": {"role": "assistant", "content": "The briefing "}, "finish_reason": None}]
),
_frame(identity, [{"index": 0, "delta": {"content": "is short."}, "finish_reason": None}]),
_frame(identity, [{"index": 0, "delta": {}, "finish_reason": "stop"}]),
_frame(identity, [], usage=_PROVIDER_USAGE),
b"data: [DONE]\n\n",
)
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/chat/completions"
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == _BACKEND
assert body["messages"] == [{"role": "user", "content": _PROMPT}]
assert body["stream"] is True
return Reply(content_type="text/event-stream", chunks=frames)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"databricks/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
with gateway.client.stream(
"POST",
"/v1/chat/completions",
json={
"model": model,
"messages": [{"role": "user", "content": _PROMPT}],
"stream": True,
"stream_options": {"include_usage": True},
},
headers={"Authorization": f"Bearer {gateway.key}"},
) as response:
assert response.status_code == 200, response.read()
lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data: "))
assert lines[-1] == "data: [DONE]", lines
chunks: Final = tuple(_Chunk.model_validate_json(line.removeprefix("data: ")) for line in lines[:-1])
assert {chunk.id for chunk in chunks} == {identity}
assert (
"".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices)
== "The briefing is short."
)
usages: Final = tuple(chunk.usage for chunk in chunks if chunk.usage is not None)
assert len(usages) == 1, lines
assert (
usages[0].prompt_tokens,
usages[0].completion_tokens,
usages[0].total_tokens,
usages[0].prompt_tokens_details.cached_tokens if usages[0].prompt_tokens_details is not None else None,
) == (12011, 8, 12019, 12002), lines
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
rows: Final = eventually(
lambda: read_rows(
'SELECT prompt_tokens, completion_tokens, total_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(identity,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert (rows[0]["prompt_tokens"], rows[0]["completion_tokens"], rows[0]["total_tokens"]) == (12011, 8, 12019)

View file

@ -0,0 +1,92 @@
import base64
import json
import uuid
from pathlib import Path
from typing import Final
from urllib.parse import parse_qs
import pytest
from integration._support.client import Gateway
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_MODEL: Final = "databricks/synthetic-vendor.chat-model.v1"
_CLIENT_ID: Final = "synthetic-databricks-client-id"
_CLIENT_SECRET: Final = "synthetic-databricks-client-secret"
_ACCESS_TOKEN: Final = "synthetic-databricks-oauth-token"
_PROMPT: Final = "Which workspace issued this token?"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
def _basic_credentials(client_id: str, client_secret: str) -> str:
return "Basic " + base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
def _completion(identity: str) -> bytes:
return json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": _MODEL.removeprefix("databricks/"),
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "the workspace origin"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 9, "completion_tokens": 4, "total_tokens": 13},
}
).encode()
@pytest.mark.covers("other.provider_wire.databricks.oauth_token_url_uses_workspace_origin_for_ai_gateway_api_base")
def test_databricks_ai_gateway_api_base_requests_oauth_token_from_workspace_origin(
gateway: Gateway, tmp_path: Path
) -> None:
identity: Final = f"databricks-oauth-{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
if request.target == "/oidc/v1/token":
assert request.method == "POST"
assert request.headers["authorization"] == _basic_credentials(_CLIENT_ID, _CLIENT_SECRET)
assert request.headers["content-type"] == "application/x-www-form-urlencoded"
assert parse_qs(request.body.decode()) == {"grant_type": ["client_credentials"], "scope": ["all-apis"]}
return Reply(
body=json.dumps({"access_token": _ACCESS_TOKEN, "token_type": "Bearer", "expires_in": 3600}).encode()
)
if request.target == "/ai-gateway/mlflow/v1/chat/completions":
assert request.method == "POST"
assert request.headers["authorization"] == f"Bearer {_ACCESS_TOKEN}"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == _MODEL.removeprefix("databricks/")
assert body["messages"] == [{"role": "user", "content": _PROMPT}]
return Reply(body=_completion(identity))
return Reply(status=401, body=json.dumps({"error": f"unauthenticated path {request.target}"}).encode())
overrides: Final = {"DATABRICKS_CLIENT_ID": _CLIENT_ID, "DATABRICKS_CLIENT_SECRET": _CLIENT_SECRET}
with wire_server(respond) as wire, owned_proxy(gateway, tmp_path, overrides) as candidate:
with candidate.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=f"{wire.url}/ai-gateway/mlflow/v1", api_key=None)
response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}]},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["id"] == identity
assert payload["choices"] == [
{
"finish_reason": "stop",
"index": 0,
"message": {"content": "the workspace origin", "role": "assistant"},
}
]
assert payload["usage"] == {"prompt_tokens": 9, "completion_tokens": 4, "total_tokens": 13}
assert [(request.method, request.target) for request in wire.drain()] == [
("POST", "/oidc/v1/token"),
("POST", "/ai-gateway/mlflow/v1/chat/completions"),
]

View file

@ -0,0 +1,40 @@
from typing import Final
import httpx
import pytest
from pydantic import JsonValue
from tests.integration._support.client import JSON_OBJECT, Gateway, object_value
_VISION_MODEL: Final = "deepseek-v4-flash-vision-exp"
_API_KEY: Final = "synthetic-deepseek-key"
_VISION_CONTENT: Final[JsonValue] = [
{"type": "text", "text": "what is in this image?"},
{"type": "image_url", "image_url": {"url": "https://example.com/pic.png"}},
]
@pytest.mark.covers("other.provider_wire.deepseek.vision_image_content_list_reaches_provider")
def test_deepseek_vision_forwards_image_url_content_list_instead_of_collapsing_to_text(gateway: Gateway) -> None:
with gateway.scenario() as scenario, httpx.Client(base_url=gateway.upstream_url, trust_env=False) as upstream:
upstream.get("/__observations").raise_for_status()
model: Final = scenario.model(
model=f"deepseek/{_VISION_MODEL}",
api_key=_API_KEY,
model_info={"mode": "chat", "supports_vision": True},
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _VISION_CONTENT}]},
)
assert response.status_code == 200, response.text
observations: Final = JSON_OBJECT.validate_json(upstream.get("/__observations").content)["requests"]
assert isinstance(observations, list)
assert len(observations) == 1, response.text
observed: Final = object_value(observations[0])
assert observed["path"] == "/v1/chat/completions", response.text
assert observed["authorization"] == f"Bearer {_API_KEY}", response.text
body: Final = object_value(observed["body"])
assert body["model"] == _VISION_MODEL, response.text
assert body["messages"] == [{"role": "user", "content": _VISION_CONTENT}], response.text

View file

@ -0,0 +1,84 @@
import json
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_ROUTER_SLUG: Final = "routers/glm-latest"
_ROUTER_RESOURCE: Final = "accounts/fireworks/routers/glm-latest"
_API_KEY: Final = "synthetic-fireworks-key"
_PROMPT: Final = "route me through the router"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
def _provider_body(request: Request, target: str) -> dict[str, JsonValue]:
assert request.method == "POST"
assert request.target == target
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
return _JSON_OBJECT.validate_json(request.body)
@pytest.mark.covers("other.provider_wire.fireworks_ai.router_slug_chat_sends_router_resource_name")
def test_fireworks_router_slug_chat_sends_router_resource_not_models_path(gateway: Gateway) -> None:
def respond(request: Request) -> Reply:
body: Final = _provider_body(request, "/chat/completions")
assert body["model"] == _ROUTER_RESOURCE, body
assert body["messages"] == [{"role": "user", "content": _PROMPT}]
return Reply(
body=json.dumps(
{
"id": "fw-router-chat",
"object": "chat.completion",
"created": 1,
"model": _ROUTER_RESOURCE,
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "routed"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6},
}
).encode()
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"fireworks_ai/{_ROUTER_SLUG}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}]},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["choices"] == [
{"finish_reason": "stop", "index": 0, "message": {"role": "assistant", "content": "routed"}}
]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
@pytest.mark.covers("other.provider_wire.fireworks_ai.router_slug_text_completion_sends_router_resource_name")
def test_fireworks_router_slug_text_completion_sends_router_resource_not_models_path(gateway: Gateway) -> None:
def respond(request: Request) -> Reply:
body: Final = _provider_body(request, "/completions")
assert body["model"] == _ROUTER_RESOURCE, body
assert body["prompt"] == _PROMPT
return Reply(
body=json.dumps(
{
"id": "fw-router-text",
"object": "text_completion",
"created": 1,
"model": _ROUTER_RESOURCE,
"choices": [{"index": 0, "text": "routed", "finish_reason": "stop", "logprobs": None}],
"usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6},
}
).encode()
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"fireworks_ai/{_ROUTER_SLUG}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/completions", {"model": model, "prompt": _PROMPT})
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["choices"] == [{"index": 0, "text": "routed", "finish_reason": "stop", "logprobs": None}]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/completions")]

View file

@ -0,0 +1,66 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "gpt-5.4-mini"
_API_KEY: Final = "synthetic-openai-key"
_PROMPT: Final = "Summarize this conversation in one sentence."
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
def _completion(identity: str, content: str) -> bytes:
return json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": _BACKEND,
"choices": [{"index": 0, "message": {"role": "assistant", "content": content}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 19, "completion_tokens": 7, "total_tokens": 26},
}
).encode()
@pytest.mark.covers("providers.openai_chat_wire.tool_choice_without_tools_is_dropped_before_the_wire")
def test_openai_chat_tool_choice_without_tools_is_not_forwarded(gateway: Gateway) -> None:
identity: Final = f"openai-toolless-{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/chat/completions"
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == _BACKEND
assert body["messages"] == [{"role": "user", "content": _PROMPT}]
assert "tool_choice" not in body, body
assert "tools" not in body, body
return Reply(body=_completion(identity, "One sentence."))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}], "tool_choice": "none"},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["id"] == identity
assert payload["choices"] == [
{
"finish_reason": "stop",
"index": 0,
"message": {
"role": "assistant",
"content": "One sentence.",
"provider_specific_fields": {"refusal": None},
},
"provider_specific_fields": {},
}
]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]

View file

@ -0,0 +1,75 @@
import json
from email.message import Message
from email.parser import BytesParser
from email.policy import HTTP
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import BaseModel
_PNG_BYTES: Final = (
b"\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01\x08\x06\x00\x00\x00"
b"\x1f\x15\xc4\x89\x00\x00\x00\rIDAT\x08\xd7c\xf8\xcf\xc0\xf0\x1f\x00\x05\x00\x01\xff"
b"\x89\x99=\x1d\x00\x00\x00\x00IEND\xaeB`\x82"
)
_PROMPT: Final = "turn the red circle green"
_EDITED_IMAGE_B64: Final = "aW50ZWdyYXRpb24tZWRpdGVkLWltYWdl"
class _Image(BaseModel):
b64_json: str
class _ImageResponse(BaseModel):
data: tuple[_Image, ...]
def _multipart_parts(request: Request) -> tuple[Message, ...]:
envelope: Final = f"content-type: {request.headers['content-type']}\r\n\r\n".encode() + request.body
parsed: Final = BytesParser(policy=HTTP).parsebytes(envelope)
assert parsed.is_multipart(), request.headers["content-type"]
return tuple(parsed.iter_parts())
def _text_fields(parts: tuple[Message, ...]) -> dict[str, str]:
return {
part.get_param("name", header="content-disposition"): part.get_payload(decode=True).decode()
for part in parts
if part.get_filename() is None
}
def _file_fields(parts: tuple[Message, ...]) -> dict[str, bytes]:
return {
part.get_param("name", header="content-disposition"): part.get_payload(decode=True)
for part in parts
if part.get_filename() is not None
}
@pytest.mark.covers("other.provider_wire.openai.image_edit_forwards_provider_specific_form_fields")
def test_openai_compatible_image_edit_forwards_seed_form_field_to_backend(gateway: Gateway) -> None:
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/v1/images/edits"
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
parts: Final = _multipart_parts(request)
assert _text_fields(parts) == {"model": "gpt-image-1", "prompt": _PROMPT, "seed": "42"}
assert _file_fields(parts) == {"image[]": _PNG_BYTES}
return Reply(body=json.dumps({"created": 1700000000, "data": [{"b64_json": _EDITED_IMAGE_B64}]}).encode())
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model="openai/gpt-image-1", api_base=f"{wire.url}/v1", api_key="synthetic-openai-key"
)
response: Final = gateway.request_multipart(
"/v1/images/edits",
{"model": model, "prompt": _PROMPT, "seed": "42"},
{"image": ("red_circle.png", _PNG_BYTES, "image/png")},
)
assert response.status_code == 200, response.text
payload: Final = _ImageResponse.model_validate_json(response.content)
assert [image.b64_json for image in payload.data] == [_EDITED_IMAGE_B64], response.text
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/v1/images/edits")]

View file

@ -0,0 +1,64 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
@pytest.mark.covers("other.provider_wire.responses_bridge.max_output_tokens_incomplete_maps_to_length")
def test_chat_over_responses_deployment_returns_length_when_output_tokens_run_out(gateway: Gateway) -> None:
identity: Final = "responses-incomplete-" + uuid.uuid4().hex
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/responses", request.target
assert request.headers["authorization"] == "Bearer synthetic-openai-key"
body: Final = json.loads(request.body)
assert body["model"] == "gpt-5.3-codex"
assert body["max_output_tokens"] == 16
assert body["reasoning"] == {"effort": "high"}
assert body["input"] == [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": f"explain the plan in detail {identity}"}],
}
]
return Reply(
body=json.dumps(
{
"id": f"resp_{identity}",
"object": "response",
"created_at": 1789788253,
"status": "incomplete",
"incomplete_details": {"reason": "max_output_tokens"},
"model": "gpt-5.3-codex",
"output": [{"type": "reasoning", "id": f"rs_{identity}", "summary": []}],
"usage": {"input_tokens": 12, "output_tokens": 16, "total_tokens": 28},
}
).encode()
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model="openai/responses/gpt-5.3-codex", api_base=wire.url, api_key="synthetic-openai-key"
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"explain the plan in detail {identity}"}],
"reasoning_effort": "high",
"max_completion_tokens": 16,
},
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert len(wire.drain()) == 1
assert [choice["finish_reason"] for choice in body["choices"]] == ["length"], response.text
assert body["choices"][0]["message"]["content"] == "", response.text
assert body["choices"][0]["message"]["role"] == "assistant", response.text
assert body["usage"]["prompt_tokens"] == 12 and body["usage"]["completion_tokens"] == 16, response.text
assert body["usage"]["total_tokens"] == 28, response.text

View file

@ -0,0 +1,86 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "deepseek-v4-pro"
_API_KEY: Final = "synthetic-tencent-key"
_PROMPT: Final = "What is 17 + 26? Answer with just the number."
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_REASONING_REQUESTS: Final[tuple[tuple[str, dict[str, JsonValue], dict[str, JsonValue]], ...]] = (
("thinking_enabled", {"thinking": {"type": "enabled"}}, {"type": "enabled"}),
("reasoning_effort_none", {"reasoning_effort": "none"}, {"type": "disabled"}),
)
def _completion(identity: str) -> bytes:
return json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": _BACKEND,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "43", "reasoning_content": "17 plus 26 is 43."},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 23, "completion_tokens": 41, "total_tokens": 64},
}
).encode()
@pytest.mark.covers("other.provider_wire.tencent.thinking_reaches_provider_in_request_body")
@pytest.mark.parametrize(
("reasoning_params", "expected_thinking"),
tuple(case[1:] for case in _REASONING_REQUESTS),
ids=tuple(case[0] for case in _REASONING_REQUESTS),
)
def test_tencent_thinking_is_sent_in_provider_body_instead_of_failing_the_request(
gateway: Gateway, reasoning_params: dict[str, JsonValue], expected_thinking: dict[str, JsonValue]
) -> None:
identity: Final = f"tencent-thinking-{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/chat/completions"
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
assert request.headers["content-type"] == "application/json"
assert _JSON_OBJECT.validate_json(request.body) == {
"model": _BACKEND,
"messages": [{"role": "user", "content": _PROMPT}],
"thinking": expected_thinking,
}
return Reply(body=_completion(identity))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"tencent/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}], **reasoning_params},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["id"] == identity
assert payload["choices"] == [
{
"finish_reason": "stop",
"index": 0,
"message": {
"role": "assistant",
"content": "43",
"reasoning_content": "17 plus 26 is 43.",
"provider_specific_fields": {"refusal": None},
},
"provider_specific_fields": {},
}
]
assert payload["usage"] == {"prompt_tokens": 23, "completion_tokens": 41, "total_tokens": 64}
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]

View file

@ -0,0 +1,299 @@
import json
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
BEDROCK_MODEL: Final = "us.anthropic.claude-haiku-4-5-20251001-v1:0"
INVOKE_TARGET: Final = f"/model/{BEDROCK_MODEL}/invoke"
SEARCH_TARGET: Final = "/tavily/search"
SEARCH_RESULT: Final = {
"title": "Synthetic result",
"url": "https://example.test/result",
"content": "the snippet text",
}
def sse_events(text: str) -> tuple[tuple[str, dict[str, object]], ...]:
frames: Final = tuple(frame for frame in text.split("\n\n") if frame.strip())
return tuple(
(
next(line.removeprefix("event: ") for line in frame.splitlines() if line.startswith("event: ")),
json.loads(next(line.removeprefix("data: ") for line in frame.splitlines() if line.startswith("data: "))),
)
for frame in frames
)
@pytest.mark.covers("other.provider_wire.bedrock.websearch_interception_streamed_capped_turn_ends_with_native_results")
def test_streamed_web_search_turn_capped_by_max_agentic_loops_ends_turn_with_snippets_and_ordered_blocks(
gateway: Gateway, tmp_path: Path
) -> None:
def respond(request: Request) -> Reply:
assert request.method == "POST", request.target
body: Final = json.loads(request.body)
if request.target == SEARCH_TARGET:
assert request.headers["authorization"] == "Bearer synthetic-tavily-key"
assert body["query"] == "query-0", body
return Reply(body=json.dumps({"query": "query-0", "results": [SEARCH_RESULT]}).encode())
assert request.target == INVOKE_TARGET
assert request.headers["authorization"] == "Bearer synthetic-bedrock-token"
assert [tool["name"] for tool in body["tools"]] == ["litellm_web_search"], body["tools"]
assert "stream" not in body, body
depth: Final = sum(
1
for message in body["messages"]
if isinstance(message["content"], list)
for block in message["content"]
if block["type"] == "tool_result"
)
if depth == 1:
assert body["messages"][2]["content"] == [
{
"type": "tool_result",
"tool_use_id": "toolu_0",
"content": "Title: Synthetic result\nURL: https://example.test/result\nSnippet: the snippet text",
}
], body["messages"]
return Reply(
body=json.dumps(
{
"id": f"msg_{depth}",
"type": "message",
"role": "assistant",
"model": BEDROCK_MODEL,
"content": [
{"type": "text", "text": f"turn-{depth}"},
{
"type": "tool_use",
"id": f"toolu_{depth}",
"name": "litellm_web_search",
"input": {"query": f"query-{depth}"},
},
],
"stop_reason": "tool_use",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 4},
}
).encode()
)
with wire_server(respond) as wire:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["search_tools"] = [
{
"search_tool_name": "integration-search",
"litellm_params": {
"search_provider": "tavily",
"api_key": "synthetic-tavily-key",
"api_base": wire.url + "/tavily",
},
}
]
config["litellm_settings"].update(
{
"callbacks": ["websearch_interception"],
"websearch_interception_params": {
"enabled_providers": ["bedrock"],
"search_tool_name": "integration-search",
"max_agentic_loops": 1,
},
}
)
path: Final = tmp_path / "websearch.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
model: Final = scenario.model(
model=f"bedrock/{BEDROCK_MODEL}",
api_key="synthetic-bedrock-token",
api_base=wire.url,
aws_region_name="us-east-1",
aws_bedrock_runtime_endpoint=wire.url,
)
response: Final = candidate.request(
"POST",
"/v1/messages",
{
"model": model,
"max_tokens": 64,
"stream": True,
"messages": [{"role": "user", "content": "search control"}],
"tools": [{"type": "web_search_20250305", "name": "web_search"}],
},
)
assert response.status_code == 200, response.text
events: Final = sse_events(response.text)
assert [name for name, _ in events][:1] == ["message_start"], response.text
assert [name for name, _ in events][-2:] == ["message_delta", "message_stop"], response.text
for position, (name, event) in enumerate(events):
if name == "content_block_stop":
assert event["index"] in {
earlier_event["index"]
for earlier, earlier_event in events[:position]
if earlier == "content_block_start"
}, response.text
started: Final = tuple(event["content_block"] for name, event in events if name == "content_block_start")
search_ids: Final = tuple(block["id"] for block in started if block["type"] == "server_tool_use")
assert search_ids and all(search_id.startswith("srvtoolu_") for search_id in search_ids), response.text
assert started[-1] == {"type": "text", "text": ""}, response.text
assert started[:-1] == tuple(
block
for search_id in search_ids
for block in (
{"type": "server_tool_use", "id": search_id, "name": "web_search", "input": {"query": "query-0"}},
{
"type": "web_search_tool_result",
"tool_use_id": search_id,
"content": [
{
"type": "web_search_result",
"url": "https://example.test/result",
"title": "Synthetic result",
"page_age": None,
"encrypted_content": "",
"snippet": "the snippet text",
}
],
},
)
), response.text
assert (
"".join(event["delta"]["text"] for name, event in events if name == "content_block_delta") == "turn-1"
), response.text
assert [event["delta"]["stop_reason"] for name, event in events if name == "message_delta"] == [
"end_turn"
], response.text
assert "litellm_web_search" not in response.text, response.text
assert [request.target for request in wire.drain()] == [INVOKE_TARGET, SEARCH_TARGET, INVOKE_TARGET]
import threading
import uuid
from typing import Final
from urllib.parse import parse_qs, urlsplit
import httpx
import pytest
from integration._support.client import Gateway, eventually
_QUERY: Final = "integration capped search"
_TEXT_BLOCK: Final = {"type": "text", "text": "searching once more"}
_NOT_INTERCEPTED: Final = "native tool reached the provider"
_SEARCH_RESULT_BLOCK: Final = {
"type": "web_search_result",
"url": "https://owned.invalid/a",
"title": "Owned result",
"page_age": None,
"encrypted_content": "",
"snippet": "owned snippet",
}
def _search_tool_use(identity: str) -> dict[str, object]:
return {"type": "tool_use", "id": identity, "name": "litellm_web_search", "input": {"query": _QUERY}}
def _anthropic_reply(identity: str, content: list[dict[str, object]], stop_reason: str) -> Reply:
return Reply(
body=json.dumps(
{
"id": identity,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": content,
"stop_reason": stop_reason,
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 4},
}
).encode()
)
@pytest.mark.covers(
"other.provider_wire.anthropic.websearch_interception_capped_loop_ends_turn_without_internal_tool_use"
)
def test_capped_websearch_interception_loop_ends_turn_instead_of_exposing_internal_tool_use(
gateway: Gateway, tmp_path: Path
) -> None:
identity: Final = "websearch-wire-" + uuid.uuid4().hex
searched: Final = threading.Event()
def respond(request: Request) -> Reply:
parts: Final = urlsplit(request.target)
if request.method == "GET" and parts.path == "/search":
assert parse_qs(parts.query)["q"] == [_QUERY], request.target
searched.set()
return Reply(
body=json.dumps(
{
"results": [
{"title": "Owned result", "url": "https://owned.invalid/a", "content": "owned snippet"}
]
}
).encode()
)
assert request.method == "POST" and parts.path == "/v1/messages", request.target
body: Final = json.loads(request.body)
if any(tool.get("type") == "web_search_20250305" for tool in body["tools"]):
return _anthropic_reply(identity, [{"type": "text", "text": _NOT_INTERCEPTED}], "end_turn")
assert [tool["name"] for tool in body["tools"]] == ["litellm_web_search"], body["tools"]
return _anthropic_reply(identity, [_TEXT_BLOCK, _search_tool_use(identity)], "tool_use")
def send(candidate: Gateway, model: str) -> httpx.Response:
return candidate.request(
"POST",
"/v1/messages",
{
"model": model,
"max_tokens": 64,
"messages": [{"role": "user", "content": identity + " attempt " + uuid.uuid4().hex}],
"tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 3}],
},
)
def searched_through_proxy(response: httpx.Response) -> bool:
return searched.is_set() and _NOT_INTERCEPTED not in response.text
with wire_server(respond) as wire:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["search_tools"] = [
{
"search_tool_name": "integration-searxng",
"litellm_params": {"search_provider": "searxng", "api_base": wire.url},
}
]
config["litellm_settings"].update(
{
"callbacks": ["websearch_interception"],
"websearch_interception_params": {
"enabled": True,
"enabled_providers": ["anthropic"],
"search_tool_name": "integration-searxng",
},
}
)
path: Final = tmp_path / "websearch.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
model: Final = scenario.model(
model="anthropic/claude-sonnet-4-5-20250929", api_base=wire.url, api_key="synthetic-anthropic-key"
)
response: Final = eventually(lambda: send(candidate, model), searched_through_proxy, seconds=40)
assert response.status_code == 200, response.text
body: Final = response.json()
assert body["stop_reason"] == "end_turn", response.text
content: Final = body["content"]
assert [block["type"] for block in content] == ["server_tool_use", "web_search_tool_result", "text"], (
response.text
)
assert content[0]["name"] == "web_search" and content[0]["input"] == {"query": _QUERY}, response.text
assert content[1]["tool_use_id"] == content[0]["id"], response.text
assert content[1]["content"] == [_SEARCH_RESULT_BLOCK], response.text
assert content[2] == _TEXT_BLOCK, response.text
targets: Final = tuple((request.method, urlsplit(request.target).path) for request in wire.drain())
assert targets[-3:] == (("POST", "/v1/messages"), ("GET", "/search"), ("POST", "/v1/messages")), targets

View file

@ -0,0 +1,84 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "grok-4.6-web-search-unmapped"
_API_KEY: Final = "synthetic-xai-key"
_SYSTEM_PROMPT: Final = "Answer in one short sentence and cite the source."
_ALLOWED_DOMAINS: Final = ("weather.example.com", "news.example.org")
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
def _responses_reply(identity: str, text: str) -> bytes:
return json.dumps(
{
"id": identity,
"object": "response",
"created_at": 1,
"status": "completed",
"model": _BACKEND,
"output": [
{
"type": "message",
"id": f"msg-{identity}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
],
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [{"type": "web_search"}],
"usage": {"input_tokens": 23, "output_tokens": 41, "total_tokens": 64},
}
).encode()
@pytest.mark.covers("other.provider_wire.xai.chat_web_search_reaches_responses_with_instructions_and_filters")
def test_xai_chat_web_search_is_sent_to_responses_with_instructions_and_nested_filters(gateway: Gateway) -> None:
identity: Final = f"xai-web-search-{uuid.uuid4().hex}"
user_prompt: Final = f"What is the weather in Paris today? Request {identity}."
def respond(request: Request) -> Reply:
assert request.method == "POST"
assert request.target == "/v1/responses", request.target
assert request.headers["authorization"] == f"Bearer {_API_KEY}"
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == _BACKEND
assert body["instructions"] == _SYSTEM_PROMPT
assert body["input"] == [
{"type": "message", "role": "user", "content": [{"type": "input_text", "text": user_prompt}]}
]
assert body["tools"] == [{"type": "web_search", "filters": {"allowed_domains": list(_ALLOWED_DOMAINS)}}]
assert "web_search_options" not in body
return Reply(body=_responses_reply(identity, "Sunny, 21C."))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"xai/{_BACKEND}", api_base=f"{wire.url}/v1", api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": user_prompt},
],
"web_search_options": {"filters": {"allowed_domains": list(_ALLOWED_DOMAINS)}},
},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
choices: Final = payload["choices"]
assert isinstance(choices, list) and len(choices) == 1, response.text
choice: Final = choices[0]
assert isinstance(choice, dict), response.text
message: Final = choice["message"]
assert isinstance(message, dict), response.text
assert (message["role"], message["content"]) == ("assistant", "Sunny, 21C."), response.text
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/v1/responses")]

View file

@ -5,6 +5,7 @@ general_settings:
store_model_in_db: true
disable_spend_logs: false
proxy_batch_write_at: 1
proxy_batch_polling_interval: 1
litellm_settings:
enable_redis_auth_cache: true
cache: true
@ -14,3 +15,11 @@ litellm_settings:
port: os.environ/REDIS_PORT
router_settings:
disable_cooldowns: true
vector_store_registry:
- vector_store_name: integration-config-store
litellm_params:
vector_store_id: vs_integration_config_store
custom_llm_provider: openai
api_base: os.environ/INTEGRATION_UPSTREAM_URL
api_key: integration-provider-key
vector_store_description: declared in tests/integration/proxy_config.yaml

View file

@ -0,0 +1,88 @@
import json
import uuid
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
ANTHROPIC_MODEL: Final = "claude-sonnet-4-5-20250929"
MODEL_RPM: Final = 40
MODEL_TPM: Final = 1000
PREMIUM_SHARE: Final = 0.5
UPSTREAM_REPLY: Final = json.dumps(
{
"id": "msg_priority_headers",
"type": "message",
"role": "assistant",
"model": ANTHROPIC_MODEL,
"content": [{"type": "text", "text": "priority header control"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 10, "output_tokens": 4},
}
).encode()
@pytest.mark.covers("other.routing.priority_rate_limits.v1_messages_success_exposes_v3_priority_headers")
def test_non_streaming_v1_messages_success_carries_v3_priority_rate_limit_headers(
gateway: Gateway, tmp_path: Path
) -> None:
probe: Final = "priority header probe " + uuid.uuid4().hex
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/messages"
assert request.headers["x-api-key"] == "synthetic-anthropic-key"
assert json.loads(request.body) == {
"model": ANTHROPIC_MODEL,
"messages": [{"role": "user", "content": probe}],
"max_tokens": 16,
"stream": False,
}
return Reply(body=UPSTREAM_REPLY)
configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
configuration["litellm_settings"] = {
**configuration["litellm_settings"],
"callbacks": ["dynamic_rate_limiter_v3"],
"priority_reservation": {"premium": PREMIUM_SHARE},
}
path: Final = tmp_path / "priority.yaml"
path.write_text(yaml.safe_dump(configuration))
with (
wire_server(respond) as wire,
owned_proxy(gateway, tmp_path, {}, config=path) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(
model=f"anthropic/{ANTHROPIC_MODEL}",
api_base=wire.url,
api_key="synthetic-anthropic-key",
rpm=MODEL_RPM,
tpm=MODEL_TPM,
)
key: Final = scenario.key(metadata={"priority": "premium"})
response: Final = candidate.request(
"POST",
"/v1/messages",
{"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": probe}]},
key=key,
)
assert response.status_code == 200, response.text
assert response.json()["content"] == [{"type": "text", "text": "priority header control"}], response.text
assert len(wire.drain()) == 1
expected: Final = {
"x-litellm-priority": "premium",
"x-litellm-rate-limiter-version": "v3",
"x-ratelimit-model_saturation_check-limit-requests": str(MODEL_RPM),
"x-ratelimit-model_saturation_check-remaining-requests": str(MODEL_RPM - 1),
"x-ratelimit-priority_model-limit-requests": str(int(MODEL_RPM * PREMIUM_SHARE)),
"x-ratelimit-priority_model-remaining-requests": str(int(MODEL_RPM * PREMIUM_SHARE) - 1),
"x-ratelimit-priority_model-limit-tokens": str(int(MODEL_TPM * PREMIUM_SHARE)),
"x-ratelimit-priority_model-remaining-tokens": str(int(MODEL_TPM * PREMIUM_SHARE) - 1),
}
observed: Final = {name: response.headers.get(name) for name in expected}
assert observed == expected, response.headers

View file

@ -0,0 +1,86 @@
import json
import threading
import uuid
from pathlib import Path
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, eventually
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
def _proxy_config(directory: Path, model: str, upstream_url: str) -> Path:
config: Final = directory / "stale_cost_map_config.yaml"
config.write_text(
json.dumps(
{
"model_list": [
{
"model_name": model,
"litellm_params": {"model": model, "api_base": upstream_url + "/v1", "api_key": "sk-upstream"},
}
],
"general_settings": {
"master_key": "os.environ/LITELLM_MASTER_KEY",
"database_url": "os.environ/DATABASE_URL",
"store_model_in_db": True,
},
"router_settings": {"disable_cooldowns": True},
}
)
)
return config
@pytest.mark.covers("other.routing.cost_map.config_deployment_dropped_by_stale_boot_map_is_restored_after_reload")
def test_config_deployment_dropped_by_stale_boot_cost_map_is_restored_after_reload(
gateway: Gateway, tmp_path: Path
) -> None:
model: Final = "integration-fresh-" + uuid.uuid4().hex
remote_map: Final = json.dumps(
{model: {"litellm_provider": "openai", "mode": "chat", "input_cost_per_token": 0, "output_cost_per_token": 0}}
).encode()
fresh_map_published: Final = threading.Event()
def respond(request: Request) -> Reply:
assert request.target == "/model_prices.json", request
return Reply(body=remote_map) if fresh_map_published.is_set() else Reply(status=503, body=b"{}")
overrides: Final = {"MODEL_COST_MAP_MIN_MODEL_COUNT": "1", "MODEL_COST_MAP_MAX_SHRINK_RATIO": "0"}
with (
wire_server(respond) as peer,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
config: Final = _proxy_config(tmp_path, model, gateway.upstream_url)
with owned_proxy(
gateway,
tmp_path,
{**overrides, "LITELLM_MODEL_COST_MAP_URL": peer.url + "/model_prices.json"},
config=config,
remove_environment=("LITELLM_LOCAL_MODEL_COST_MAP",),
) as candidate:
assert model not in tuple(entry["id"] for entry in candidate.get("/v1/models")["data"])
fresh_map_published.set()
reload: Final = candidate.request("POST", "/reload/model_cost_map")
assert reload.status_code == 200, reload.text
eventually(
lambda: tuple(str(entry["id"]) for entry in candidate.get("/v1/models")["data"]),
lambda served: model in served,
seconds=30,
)
upstream.get("/__observations").raise_for_status()
response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "stale cost map control"}]},
)
assert response.status_code == 200, response.text
assert upstream.get("/__observations").json()["requests"] == [
{
"path": "/v1/chat/completions",
"authorization": "Bearer sk-upstream",
"body": {"model": model, "messages": [{"role": "user", "content": "stale cost map control"}]},
}
]

View file

@ -47,6 +47,7 @@ def main() -> int:
"pytest",
*selected,
"-vv",
"-rs",
"--strict-markers",
"-p",
"no:pytest-retry",
@ -69,7 +70,8 @@ def main() -> int:
if result != 0:
return result
evidence: Final = json.loads((output / "execution.json").read_text())
if not evidence["complete"] or sorted(evidence["passed"]) != expected or sorted(evidence["collected"]) != expected:
executed: Final = sorted(evidence["passed"] + evidence["skipped"])
if not evidence["complete"] or executed != expected or sorted(evidence["collected"]) != expected:
print("Executed integration nodes differ from the canonical manifest", file=sys.stderr)
return 1
return 0

View file

@ -0,0 +1,202 @@
from __future__ import annotations
import json
import uuid
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import JSON_OBJECT, Gateway, eventually, string_value
from integration._support.database import read_rows
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
from pydantic import JsonValue
FIRST_LINE: Final = {"prompt_tokens": 10, "completion_tokens": 7, "reasoning_tokens": 4}
SECOND_LINE: Final = {"prompt_tokens": 5, "completion_tokens": 3, "reasoning_tokens": 2}
ERROR_FILE_LINES: Final = 2
def _succeeded_line(index: int, model: str, prompt_tokens: int, completion_tokens: int, reasoning_tokens: int) -> str:
return json.dumps(
{
"id": f"batch_req_{index}",
"custom_id": f"r{index}",
"response": {
"status_code": 200,
"request_id": f"$REQUEST_ID-{index}",
"body": {
"id": f"chatcmpl-$REQUEST_ID-{index}",
"object": "chat.completion",
"model": model,
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}
],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
"completion_tokens_details": {"reasoning_tokens": reasoning_tokens},
},
},
},
"error": None,
},
separators=(",", ":"),
)
def _failed_line(index: int) -> str:
return json.dumps(
{
"id": f"batch_req_{index}",
"custom_id": f"r{index}",
"response": {
"status_code": 400,
"request_id": f"$REQUEST_ID-{index}",
"body": {"error": {"message": "rejected line", "type": "invalid_request_error", "code": "400"}},
},
"error": {"code": "bad_request", "message": "rejected line"},
},
separators=(",", ":"),
)
def _batch_routes(model: str) -> RoutedResponse:
output_lines: Final = (
_succeeded_line(1, model, **FIRST_LINE),
_succeeded_line(2, model, **SECOND_LINE),
_failed_line(3),
)
error_lines: Final = tuple(_failed_line(index) for index in range(4, 4 + ERROR_FILE_LINES))
completed: Final = {
"id": "batch-$REQUEST_ID",
"object": "batch",
"endpoint": "/v1/chat/completions",
"errors": None,
"input_file_id": "file-in-$REQUEST_ID",
"completion_window": "24h",
"status": "completed",
"output_file_id": "file-out-$REQUEST_ID",
"error_file_id": "file-err-$REQUEST_ID",
"created_at": 1,
"in_progress_at": 1,
"completed_at": 1,
"expires_at": 1,
"request_counts": {"total": 5, "completed": 2, "failed": 3},
"metadata": None,
}
return RoutedResponse(
content_type="application/x-routed",
routes={
"POST /files": JsonResponse(
content_type="application/json",
body={
"id": "file-in-$REQUEST_ID",
"object": "file",
"purpose": "batch",
"bytes": 100,
"created_at": 1,
"filename": "in.jsonl",
"status": "processed",
},
),
"POST /batches": JsonResponse(
content_type="application/json",
body={**completed, "status": "validating", "output_file_id": None, "error_file_id": None},
),
"GET /batches/batch-$REQUEST_ID": JsonResponse(content_type="application/json", body=completed),
"GET /files/file-out-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body="\n".join(output_lines) + "\n"
),
"GET /files/file-err-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body="\n".join(error_lines) + "\n"
),
},
)
def _input_file(model: str) -> bytes:
return (
"\n".join(
json.dumps(
{
"custom_id": f"r{index}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": model, "messages": [{"role": "user", "content": "batch accounting"}]},
},
separators=(",", ":"),
)
for index in range(1, 6)
)
+ "\n"
).encode()
def _metadata(value: object) -> dict[str, JsonValue]:
return JSON_OBJECT.validate_json(value) if isinstance(value, str) else JSON_OBJECT.validate_python(value)
@pytest.mark.covers("quota_management.spend_tracking.batch_costs.reasoning_tokens_and_error_file_failures_recorded")
def test_completed_batch_spend_row_records_reasoning_tokens_and_error_file_failures(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key()
scenario_id: Final = f"batch-accounting-{uuid.uuid4().hex[:12]}"
handle: Final = register_scenario(scenario_id, _batch_routes("gpt-4o-mini"))
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(api_base=handle.api_base())
file_response: Final = gateway.request_multipart(
"/v1/files",
{"purpose": "batch", "model": model},
{"file": ("in.jsonl", _input_file(model), "application/jsonl")},
key=key,
)
assert file_response.status_code == 200, file_response.text
batch_response: Final = gateway.request(
"POST",
"/v1/batches",
{
"input_file_id": string_value(JSON_OBJECT.validate_json(file_response.content)["id"]),
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"model": model,
},
key=key,
)
assert batch_response.status_code == 200, batch_response.text
batch_id: Final = string_value(JSON_OBJECT.validate_json(batch_response.content)["id"])
retrieval: Final = gateway.request("GET", f"/v1/batches/{batch_id}", key=key)
assert retrieval.status_code == 200, retrieval.text
assert retrieval.json()["status"] == "completed", retrieval.text
rows: Final = eventually(
lambda: read_rows(
'SELECT status, prompt_tokens, completion_tokens, metadata FROM "LiteLLM_SpendLogs" '
"WHERE api_key=%s AND call_type='aretrieve_batch'",
(sha256(key.encode()).hexdigest(),),
),
lambda values: len(values) == 1,
seconds=70,
)
row: Final = rows[0]
metadata: Final = _metadata(row["metadata"])
prompt_tokens: Final = FIRST_LINE["prompt_tokens"] + SECOND_LINE["prompt_tokens"]
completion_tokens: Final = FIRST_LINE["completion_tokens"] + SECOND_LINE["completion_tokens"]
reasoning_tokens: Final = FIRST_LINE["reasoning_tokens"] + SECOND_LINE["reasoning_tokens"]
assert row["status"] == "success", retrieval.text
assert (row["prompt_tokens"], row["completion_tokens"]) == (prompt_tokens, completion_tokens), retrieval.text
assert (metadata["batch_successful_requests"], metadata["batch_failed_requests"]) == (
2,
1 + ERROR_FILE_LINES,
), json.dumps(metadata)
usage: Final = JSON_OBJECT.validate_python(metadata["usage_object"])
details: Final = JSON_OBJECT.validate_python(usage["completion_tokens_details"])
assert (usage["prompt_tokens"], usage["completion_tokens"], usage["total_tokens"]) == (
prompt_tokens,
completion_tokens,
prompt_tokens + completion_tokens,
), json.dumps(metadata)
assert {name: value for name, value in details.items() if value is not None} == {
"reasoning_tokens": reasoning_tokens,
"text_tokens": completion_tokens - reasoning_tokens,
}, json.dumps(metadata)

View file

@ -0,0 +1,201 @@
from __future__ import annotations
import json
import uuid
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
from pydantic import JsonValue
REASONING_TOKENS: Final = (30, 50)
PROMPT_TOKENS: Final = 10
COMPLETION_TOKENS: Final = 100
ERROR_FILE_FAILURES: Final = 2
def _successful_line(index: int, reasoning_tokens: int) -> str:
return json.dumps(
{
"id": f"batch_req_{index}",
"custom_id": f"r{index}",
"response": {
"status_code": 200,
"request_id": f"$REQUEST_ID-{index}",
"body": {
"id": f"chatcmpl-$REQUEST_ID-{index}",
"object": "chat.completion",
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}
],
"usage": {
"prompt_tokens": PROMPT_TOKENS,
"completion_tokens": COMPLETION_TOKENS,
"total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS,
"completion_tokens_details": {"reasoning_tokens": reasoning_tokens},
},
},
},
"error": None,
},
separators=(",", ":"),
)
def _failed_line(index: int) -> str:
return json.dumps(
{
"id": f"batch_req_{index}",
"custom_id": f"r{index}",
"response": {"status_code": 400, "request_id": f"$REQUEST_ID-{index}", "body": {"error": "bad"}},
"error": {"code": "bad_request", "message": "failed"},
},
separators=(",", ":"),
)
def _batch(status: str, *, files_ready: bool) -> dict[str, JsonValue]:
return {
"id": "batch-$REQUEST_ID",
"object": "batch",
"endpoint": "/v1/chat/completions",
"errors": None,
"input_file_id": "file-in-$REQUEST_ID",
"completion_window": "24h",
"status": status,
"output_file_id": "file-out-$REQUEST_ID" if files_ready else None,
"error_file_id": "file-err-$REQUEST_ID" if files_ready else None,
"created_at": 1,
"in_progress_at": 1,
"completed_at": 1 if files_ready else None,
"expires_at": 1,
"request_counts": {"total": 5, "completed": 2, "failed": 3},
"metadata": None,
}
def _provider_routes() -> RoutedResponse:
output_lines: Final = (
_successful_line(1, REASONING_TOKENS[0]),
_failed_line(2),
_successful_line(3, REASONING_TOKENS[1]),
)
error_lines: Final = tuple(_failed_line(index) for index in range(4, 4 + ERROR_FILE_FAILURES))
return RoutedResponse(
content_type="application/x-routed",
routes={
"POST /files": JsonResponse(
content_type="application/json",
body={
"id": "file-in-$REQUEST_ID",
"object": "file",
"purpose": "batch",
"bytes": 100,
"created_at": 1,
"filename": "in.jsonl",
"status": "processed",
},
),
"POST /batches": JsonResponse(
content_type="application/json", body=_batch("validating", files_ready=False)
),
"GET /batches/batch-$REQUEST_ID": JsonResponse(
content_type="application/json", body=_batch("completed", files_ready=True)
),
"GET /files/file-out-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body="\n".join(output_lines) + "\n"
),
"GET /files/file-err-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body="\n".join(error_lines) + "\n"
),
},
)
def _input_file(model_name: str) -> bytes:
return (
"\n".join(
json.dumps(
{
"custom_id": f"r{index}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": model_name, "messages": [{"role": "user", "content": "batch observability"}]},
},
separators=(",", ":"),
)
for index in range(1, 6)
)
+ "\n"
).encode()
def _retrieval_rows(key: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(
read_rows(
'SELECT prompt_tokens, completion_tokens, metadata FROM "LiteLLM_SpendLogs" '
"WHERE api_key=%s AND call_type='aretrieve_batch'",
(sha256(key.encode()).hexdigest(),),
)
)
def _metadata(row: dict[str, JsonValue]) -> dict[str, JsonValue]:
value: Final = row["metadata"]
return object_value(JSON_OBJECT.validate_json(value) if isinstance(value, str) else value)
@pytest.mark.covers("spend.batches.retrieval_row_aggregates_reasoning_tokens_and_per_request_counts")
def test_batch_retrieval_row_sums_reasoning_tokens_and_counts_output_and_error_file_failures(
gateway: Gateway,
) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"batch-observability-{uuid.uuid4().hex[:12]}"
handle: Final = register_scenario(scenario_id, _provider_routes())
scenario.cleanups.callback(delete_scenario, handle)
model_name: Final = scenario.model(api_base=handle.api_base())
key: Final = scenario.key(models=[model_name])
file_response: Final = gateway.request_multipart(
"/v1/files",
{"purpose": "batch", "model": model_name},
{"file": ("in.jsonl", _input_file(model_name), "application/jsonl")},
key=key,
)
assert file_response.status_code == 200, file_response.text
input_file_id: Final = string_value(JSON_OBJECT.validate_json(file_response.content)["id"])
batch_response: Final = gateway.request(
"POST",
"/v1/batches",
{
"input_file_id": input_file_id,
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"model": model_name,
},
key=key,
)
assert batch_response.status_code == 200, batch_response.text
batch_id: Final = string_value(JSON_OBJECT.validate_json(batch_response.content)["id"])
retrieval: Final = eventually(
lambda: gateway.request("GET", f"/v1/batches/{batch_id}", key=key),
lambda response: response.status_code == 200 and response.json()["status"] == "completed",
seconds=30,
)
assert retrieval.status_code == 200, retrieval.text
rows: Final = eventually(lambda: _retrieval_rows(key), lambda values: len(values) == 1, seconds=70)
row: Final = rows[0]
metadata: Final = _metadata(row)
usage: Final = object_value(metadata["usage_object"])
assert row["prompt_tokens"] == 2 * PROMPT_TOKENS, retrieval.text
assert row["completion_tokens"] == 2 * COMPLETION_TOKENS, retrieval.text
assert object_value(usage["completion_tokens_details"])["reasoning_tokens"] == sum(REASONING_TOKENS), (
retrieval.text,
usage,
)
assert metadata["batch_successful_requests"] == 2, (retrieval.text, metadata)
assert metadata["batch_failed_requests"] == 1 + ERROR_FILE_FAILURES, (retrieval.text, metadata)

View file

@ -0,0 +1,212 @@
import json
import os
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse, TextResponse
from pydantic import JsonValue
from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE
INPUT_COST_PER_TOKEN: Final = 0.001
OUTPUT_COST_PER_TOKEN: Final = 0.002
PROMPT_TOKENS: Final = 100
COMPLETION_TOKENS: Final = 50
BATCH_COST_SHARE: Final = 0.5
_INPUT_FILE: Final = JsonResponse(
content_type="application/json",
body={
"id": "file-in-$REQUEST_ID",
"object": "file",
"purpose": "batch",
"bytes": 100,
"created_at": 1,
"filename": "in.jsonl",
"status": "processed",
},
)
def _batch(status: str, output_file_id: str | None) -> dict[str, JsonValue]:
return {
"id": "batch-$REQUEST_ID",
"object": "batch",
"endpoint": "/v1/chat/completions",
"errors": None,
"input_file_id": "file-in-$REQUEST_ID",
"completion_window": "24h",
"status": status,
"output_file_id": output_file_id,
"error_file_id": None,
"created_at": 1,
"in_progress_at": 1,
"completed_at": 1 if status == "completed" else None,
"expires_at": 1,
"request_counts": {"total": 1, "completed": 1 if status == "completed" else 0, "failed": 0},
"metadata": None,
}
def _accepting_routes() -> dict[str, JsonResponse | TextResponse]:
return {
"POST /files": _INPUT_FILE,
"POST /batches": JsonResponse(content_type="application/json", body=_batch("validating", None)),
}
def _gone_at_provider_routes() -> RoutedResponse:
return RoutedResponse(
content_type="application/x-routed",
routes={
**_accepting_routes(),
"GET /batches/batch-$REQUEST_ID": JsonResponse(
content_type="application/json",
status=404,
body={
"error": {
"message": "No batch found with id 'batch-$REQUEST_ID'.",
"type": "invalid_request_error",
"param": "id",
"code": "batch_not_found",
}
},
),
},
)
def _completed_routes() -> RoutedResponse:
output_line: Final = {
"id": "batch_req_1",
"custom_id": "r1",
"response": {
"status_code": 200,
"request_id": "$REQUEST_ID-1",
"body": {
"id": "chatcmpl-$REQUEST_ID-1",
"object": "chat.completion",
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {
"prompt_tokens": PROMPT_TOKENS,
"completion_tokens": COMPLETION_TOKENS,
"total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS,
},
},
},
"error": None,
}
return RoutedResponse(
content_type="application/x-routed",
routes={
**_accepting_routes(),
"GET /batches/batch-$REQUEST_ID": JsonResponse(
content_type="application/json", body=_batch("completed", "file-out-$REQUEST_ID")
),
"GET /files/file-out-$REQUEST_ID/content": TextResponse(
content_type="application/jsonl", body=json.dumps(output_line, separators=(",", ":")) + "\n"
),
},
)
def _scripted_deployment(scenario: Scenario, marker: str, routes: RoutedResponse) -> str:
scenario_id: Final = f"poll-{marker}-{sha256(os.urandom(16)).hexdigest()[:12]}"
handle: Final = register_scenario(scenario_id, routes)
scenario.cleanups.callback(delete_scenario, handle)
created: Final = scenario.gateway.post(
"/model/new",
{
"model_name": f"poll-{marker}-{sha256(scenario_id.encode()).hexdigest()[:12]}",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-scripted-provider",
"api_base": handle.api_base(),
"input_cost_per_token": INPUT_COST_PER_TOKEN,
"output_cost_per_token": OUTPUT_COST_PER_TOKEN,
},
},
)
scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"]))
return string_value(created["model_name"])
def _submitted_batch_id(gateway: Gateway, key: str, model_name: str) -> str:
request_line: Final = {
"custom_id": "r1",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": model_name, "messages": [{"role": "user", "content": "poll starvation"}]},
}
file_response: Final = gateway.request_multipart(
"/v1/files",
{"purpose": "batch", "target_model_names": model_name},
{"file": ("in.jsonl", (json.dumps(request_line) + "\n").encode(), "application/jsonl")},
key=key,
)
assert file_response.is_success, file_response.text
batch_response: Final = gateway.request(
"POST",
"/v1/batches",
{
"input_file_id": string_value(JSON_OBJECT.validate_json(file_response.content)["id"]),
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
"model": model_name,
},
key=key,
)
assert batch_response.is_success, batch_response.text
return string_value(JSON_OBJECT.validate_json(batch_response.content)["id"])
def _managed_rows(batch_ids: tuple[str, ...]) -> list[dict[str, JsonValue]]:
placeholders: Final = ", ".join("%s" for _ in batch_ids)
return read_rows(
f'SELECT batch_processed FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id IN ({placeholders})',
batch_ids,
)
@pytest.mark.timeout(180)
@pytest.mark.covers("quota_management.spend_tracking.batch_costs.uncostable_rows_retire_so_newer_batches_are_costed")
def test_batches_gone_at_provider_do_not_starve_a_newer_batch_out_of_cost_polling(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key()
gone_batch_ids: Final = tuple(
_submitted_batch_id(
gateway, key, _scripted_deployment(scenario, f"gone{index}", _gone_at_provider_routes())
)
for index in range(MAX_OBJECTS_PER_POLL_CYCLE)
)
costable_batch_id: Final = _submitted_batch_id(
gateway, key, _scripted_deployment(scenario, "costable", _completed_routes())
)
spend_rows: Final = eventually(
lambda: read_rows(
'SELECT call_type, status, prompt_tokens, completion_tokens, spend FROM "LiteLLM_SpendLogs" '
"WHERE api_key = %s AND call_type = %s",
(sha256(key.encode()).hexdigest(), "aretrieve_batch"),
),
lambda rows: len(rows) == 1,
seconds=120,
)
assert spend_rows == [
{
"call_type": "aretrieve_batch",
"status": "success",
"prompt_tokens": PROMPT_TOKENS,
"completion_tokens": COMPLETION_TOKENS,
"spend": pytest.approx(
BATCH_COST_SHARE
* (PROMPT_TOKENS * INPUT_COST_PER_TOKEN + COMPLETION_TOKENS * OUTPUT_COST_PER_TOKEN)
),
}
]
assert _managed_rows((costable_batch_id,)) == [{"batch_processed": True}]
assert _managed_rows(gone_batch_ids) == [{"batch_processed": True}] * MAX_OBJECTS_PER_POLL_CYCLE

View file

@ -1,4 +1,7 @@
import json
import threading
import uuid
from concurrent.futures import ThreadPoolExecutor
from contextlib import ExitStack
from hashlib import sha256
from typing import Final
@ -7,10 +10,10 @@ import httpx
import pytest
from hypothesis import strategies as st
from hypothesis.stateful import RuleBasedStateMachine, rule, run_state_machine_as_test
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests
from integration._support.wire import Reply, Request, wire_server
@pytest.mark.covers("quota_management.response_cache.generated_sequences_preserve_content_and_accounting")
@ -211,6 +214,117 @@ def test_key_budget_at_boundary_blocks_provider_then_explicit_reset_restores(gat
assert upstream.get("/__observations").json()["requests"] == []
@pytest.mark.covers("quota_management.budget.key.count_tokens_reserves_nothing_so_completion_within_budget_succeeds")
def test_repeated_count_tokens_on_budgeted_key_does_not_reserve_budget_or_block_later_completion(
gateway: Gateway,
) -> None:
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
key: Final = scenario.key(models=[model], max_budget=0.1)
digest: Final = sha256(key.encode()).hexdigest()
upstream.get("/__observations").raise_for_status()
counts: Final = tuple(
gateway.request(
"POST",
"/v1/messages/count_tokens",
{"model": model, "messages": [{"role": "user", "content": "hello!!!"}]},
key=key,
headers={"anthropic-version": "2023-06-01"},
)
for _ in range(3)
)
for count in counts:
assert count.status_code == 200, count.text
assert count.json() == counts[0].json(), count.text
input_tokens: Final = counts[0].json()["input_tokens"]
assert isinstance(input_tokens, int) and input_tokens > 0, counts[0].text
assert upstream.get("/__observations").json()["requests"] == []
completion: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"after counting {uuid.uuid4().hex}"}]},
key=key,
)
assert completion.status_code == 200, completion.text
assert completion.json()["usage"]["total_tokens"] == 40, completion.text
assert [request["path"] for request in upstream.get("/__observations").json()["requests"]] == [
"/v1/chat/completions"
]
spent: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)),
lambda values: len(values) == 1 and float(values[0]["spend"]) > 0,
seconds=70,
)
assert float(spent[0]["spend"]) == pytest.approx(20 * 0.001 + 20 * 0.002)
rows: Final = eventually(
lambda: read_rows('SELECT call_type, spend FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)),
lambda values: len(values) >= 1,
seconds=70,
)
assert [(row["call_type"], float(row["spend"])) for row in rows] == [("acompletion", pytest.approx(0.06))]
@pytest.mark.covers(
"quota_management.budget.key.in_flight_count_tokens_reserves_nothing_so_completion_reaches_provider"
)
def test_in_flight_count_tokens_does_not_reserve_key_budget_away_from_a_completion(gateway: Gateway) -> None:
counting_reached_provider: Final = threading.Event()
completion_answered: Final = threading.Event()
def respond(request: Request) -> Reply:
counting_reached_provider.set()
assert completion_answered.wait(timeout=30), "completion never ran while count tokens was in flight"
return Reply(body=b'{"totalTokens": 12, "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 12}]}')
with (
wire_server(respond) as wire,
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
ThreadPoolExecutor(max_workers=1) as background,
):
counted: Final = scenario.model(
model="gemini/gemini-3.8-flash",
api_base=wire.url,
api_key="synthetic-gemini-key",
input_cost_per_token=0.001,
output_cost_per_token=0.002,
)
completed: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
key: Final = scenario.key(models=[counted, completed], max_budget=0.06)
contents: Final = [{"role": "user", "parts": [{"text": "hello"}]}]
counting: Final = background.submit(
gateway.request, "POST", f"/v1beta/models/{counted}:countTokens", {"contents": contents}, key=key
)
assert counting_reached_provider.wait(timeout=30), "count tokens request never reached the provider"
upstream.get("/__observations").raise_for_status()
prompt: Final = f"after count tokens {uuid.uuid4().hex}"
completion: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": completed, "messages": [{"role": "user", "content": prompt}]},
key=key,
)
completion_answered.set()
count: Final = counting.result(timeout=30)
assert completion.status_code == 200 and completion.json()["usage"]["total_tokens"] == 40, completion.text
assert [call["body"]["messages"] for call in upstream.get("/__observations").json()["requests"]] == [
[{"role": "user", "content": prompt}]
]
assert count.status_code == 200, count.text
assert count.json() == {"totalTokens": 12, "promptTokensDetails": [{"modality": "TEXT", "tokenCount": 12}]}, (
count.text
)
provider_calls: Final = wire.drain()
assert [(call.method, call.target) for call in provider_calls] == [
("POST", "/v1beta/models/gemini-3.8-flash:countTokens")
]
assert provider_calls[0].headers["x-goog-api-key"] == "synthetic-gemini-key"
assert json.loads(provider_calls[0].body) == {"contents": contents}
@pytest.mark.covers("quota_management.response_cache.system_messages_partition_cache_identity")
def test_different_system_messages_do_not_share_a_cached_response(gateway: Gateway) -> None:
with (
@ -219,8 +333,7 @@ def test_different_system_messages_do_not_share_a_cached_response(gateway: Gatew
):
model: Final = scenario.model()
prompt: Final = uuid.uuid4().hex
identities: dict[str, str] = {}
for system, expected_calls in (("first policy", 1), ("second policy", 1), ("first policy", 0)):
def completion_id(system: str, expected_calls: int) -> str:
upstream.get("/__observations").raise_for_status()
response: Final = gateway.request(
"POST",
@ -232,14 +345,12 @@ def test_different_system_messages_do_not_share_a_cached_response(gateway: Gatew
)
assert response.status_code == 200 and response.json()["usage"]["total_tokens"] == 40, response.text
calls: Final = upstream.get("/__observations").json()["requests"]
assert len(calls) == expected_calls
if system in identities:
assert response.json()["id"] == identities[system]
else:
assert response.json()["id"] not in identities.values()
identities = {**identities, system: response.json()["id"]}
if calls:
assert calls[0]["body"]["messages"] == [
{"role": "system", "content": system},
{"role": "user", "content": prompt},
]
assert [call["body"]["messages"] for call in calls] == [
[{"role": "system", "content": system}, {"role": "user", "content": prompt}]
] * expected_calls, calls
return response.json()["id"]
first_policy_id: Final = completion_id("first policy", 1)
second_policy_id: Final = completion_id("second policy", 1)
assert first_policy_id != second_policy_id
assert completion_id("first policy", 0) == first_policy_id

View file

@ -0,0 +1,182 @@
import json
import os
import uuid
from collections.abc import Iterable
from hashlib import sha256
from typing import Final
import psycopg
import pytest
from integration._support.client import Gateway, delete_key_if_present, eventually, string_value
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
from psycopg import sql
def _execute(statements: Iterable[sql.Composable]) -> None:
with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
for statement in statements:
connection.execute(statement)
def _install_daily_user_rollup_fault(user_id: str) -> str:
suffix: Final = f"fault-{uuid.uuid4().hex}"
sequence: Final = sql.Identifier(f"{suffix}_attempts")
function: Final = sql.Identifier(suffix)
_execute(
(
sql.SQL("CREATE SEQUENCE {}").format(sequence),
sql.SQL(
"CREATE FUNCTION {}() RETURNS trigger LANGUAGE plpgsql AS $fault$ "
"BEGIN PERFORM nextval({}); "
"RAISE EXCEPTION 'synthetic daily rollup outage' USING ERRCODE = '55P03'; "
"END $fault$"
).format(function, sql.Literal(f"{suffix}_attempts")),
sql.SQL(
'CREATE TRIGGER {} BEFORE INSERT ON "LiteLLM_DailyUserSpend" '
"FOR EACH ROW WHEN (NEW.user_id = {}) EXECUTE FUNCTION {}()"
).format(sql.Identifier(suffix), sql.Literal(user_id), function),
)
)
return suffix
def _lift_daily_user_rollup_fault(suffix: str) -> None:
_execute(
(
sql.SQL('DROP TRIGGER IF EXISTS {} ON "LiteLLM_DailyUserSpend"').format(sql.Identifier(suffix)),
sql.SQL("DROP FUNCTION IF EXISTS {}()").format(sql.Identifier(suffix)),
sql.SQL("DROP SEQUENCE IF EXISTS {}").format(sql.Identifier(f"{suffix}_attempts")),
)
)
def _rollup_attempts(suffix: str) -> int:
with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
row: Final = connection.execute(
sql.SQL("SELECT CASE WHEN is_called THEN last_value ELSE 0 END FROM {}").format(
sql.Identifier(f"{suffix}_attempts")
)
).fetchone()
assert row is not None
return int(row[0])
@pytest.mark.covers("spend.daily_rollup.failed_user_commit_is_retried_until_report_and_daily_activity_agree")
def test_failed_daily_user_rollup_commit_is_retried_so_spend_report_and_daily_activity_agree(
gateway: Gateway,
) -> None:
def provider(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/chat/completions"
return Reply(
body=json.dumps(
{
"id": "chatcmpl-" + uuid.uuid4().hex,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "synthetic rollup answer"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 20, "completion_tokens": 20, "total_tokens": 40},
}
).encode()
)
with wire_server(provider) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
api_base=wire.url + "/v1", input_cost_per_token=0.001, output_cost_per_token=0.002, num_retries=0
)
user: Final = scenario.user()
key: Final = string_value(gateway.post("/key/generate", {"user_id": user, "models": [model]})["key"])
scenario.cleanups.callback(delete_key_if_present, gateway, key)
digest: Final = sha256(key.encode()).hexdigest()
suffix: Final = _install_daily_user_rollup_fault(user)
scenario.cleanups.callback(_lift_daily_user_rollup_fault, suffix)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "rollup retry control"}]},
key=key,
)
assert response.status_code == 200, response.text
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(20 * 0.001 + 20 * 0.002)
body: Final = response.json()
spend_rows: Final = eventually(
lambda: read_rows(
'SELECT spend, DATE("startTime")::text AS day, model FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(body["id"],),
),
lambda values: len(values) == 1,
seconds=70,
)
assert float(spend_rows[0]["spend"]) == pytest.approx(0.06)
day: Final = string_value(spend_rows[0]["day"])
stored_model: Final = string_value(spend_rows[0]["model"])
eventually(lambda: _rollup_attempts(suffix), lambda attempts: attempts >= 1, seconds=70)
_lift_daily_user_rollup_fault(suffix)
activity: Final = eventually(
lambda: gateway.request(
"GET",
"/user/daily/activity/aggregated",
params={"start_date": day, "end_date": day, "api_key": digest},
),
lambda polled: (
polled.status_code == 200
and len(polled.json().get("results", ())) > 0
and polled.json()["results"][0]["breakdown"]["api_keys"]
.get(digest, {})
.get("metrics", {})
.get("spend", 0)
== pytest.approx(0.06)
),
seconds=90,
)
assert activity.status_code == 200, activity.text
metrics: Final = activity.json()["results"][0]["breakdown"]["api_keys"][digest]["metrics"]
assert metrics == {
"spend": pytest.approx(0.06),
"flat_cost": pytest.approx(0.0),
"prompt_tokens": 20,
"completion_tokens": 20,
"cache_read_input_tokens": 0,
"cache_creation_input_tokens": 0,
"compression_saved_tokens": 0,
"compression_savings_spend": pytest.approx(0.0),
"prompt_caching_savings_spend": pytest.approx(0.0),
"gateway_injected_caching_savings_spend": pytest.approx(0.0),
"autorouter_savings_spend": pytest.approx(0.0),
"total_tokens": 40,
"successful_requests": 1,
"failed_requests": 0,
"api_requests": 1,
"total_response_time_ms": metrics["total_response_time_ms"],
"timed_requests": metrics["timed_requests"],
}
report: Final = gateway.request(
"GET",
"/global/spend/report",
params={"start_date": day, "end_date": day, "api_key": digest},
)
assert report.status_code == 200, report.text
assert report.json() == [
{
"api_key": digest,
"total_cost": pytest.approx(0.06),
"total_input_tokens": 20,
"total_output_tokens": 20,
"model_details": [
{
"model": stored_model,
"total_cost": pytest.approx(0.06),
"total_input_tokens": 20,
"total_output_tokens": 20,
}
],
}
]
assert metrics["spend"] == pytest.approx(report.json()[0]["total_cost"])

View file

@ -0,0 +1,123 @@
import base64
import json
import uuid
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.upstream import _aws_event_frame
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue
BEDROCK_MODEL: Final = "anthropic.claude-haiku-4-5-20251001-v1:0"
INPUT_TOKENS: Final = 30
FULL_OUTPUT_TOKENS: Final = 412
INPUT_RATE: Final = 0.001
OUTPUT_RATE: Final = 0.002
def _invoke_chunk(payload: dict[str, JsonValue]) -> bytes:
encoded: Final = base64.b64encode(json.dumps(payload, separators=(",", ":")).encode()).decode()
return _aws_event_frame("chunk", {"bytes": encoded}, "", "")
def _message_start(message_id: str) -> bytes:
return _invoke_chunk(
{
"type": "message_start",
"message": {
"id": message_id,
"type": "message",
"role": "assistant",
"model": BEDROCK_MODEL,
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": INPUT_TOKENS, "output_tokens": 0},
},
}
) + _invoke_chunk({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}})
def _text_delta(text: str) -> bytes:
return _invoke_chunk({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}})
def _terminal_usage() -> bytes:
return (
_invoke_chunk({"type": "content_block_stop", "index": 0})
+ _invoke_chunk(
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn"},
"usage": {"output_tokens": FULL_OUTPUT_TOKENS},
}
)
+ _invoke_chunk({"type": "message_stop"})
)
@pytest.mark.covers("spend.anthropic_messages_stream.client_disconnect_bills_terminal_bedrock_usage")
@pytest.mark.timeout(120)
def test_client_disconnect_mid_bedrock_messages_stream_still_bills_terminal_usage(gateway: Gateway) -> None:
message_id: Final = f"msg_{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
assert request.target == f"/model/{BEDROCK_MODEL}/invoke-with-response-stream", request.target
assert json.loads(request.body)["messages"] == [{"role": "user", "content": "disconnect control"}], request.body
return Reply(
content_type="application/vnd.amazon.eventstream",
chunks=(
_message_start(message_id) + _text_delta("first"),
_text_delta("second"),
_text_delta("third"),
_terminal_usage(),
),
pause_between_chunks=0.5,
)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"bedrock/invoke/{BEDROCK_MODEL}",
api_base=wire.url,
aws_access_key_id="AKIASCRIPTEDPROVIDER",
aws_secret_access_key="scripted-secret",
aws_region_name="us-east-1",
input_cost_per_token=INPUT_RATE,
output_cost_per_token=OUTPUT_RATE,
)
key: Final = scenario.key(models=[model])
with gateway.client.stream(
"POST",
"/v1/messages",
json={
"model": model,
"messages": [{"role": "user", "content": "disconnect control"}],
"max_tokens": FULL_OUTPUT_TOKENS,
"stream": True,
},
headers={"Authorization": f"Bearer {key}"},
) as response:
assert response.status_code == 200, response.read().decode()
first_event: Final = next(line for line in response.iter_lines() if line.startswith("data:"))
assert json.loads(first_event.removeprefix("data:"))["type"] == "message_start", first_event
rows: Final = eventually(
lambda: read_rows(
'SELECT request_id, status, prompt_tokens, completion_tokens, spend FROM "LiteLLM_SpendLogs" '
"WHERE api_key=%s",
(sha256(key.encode()).hexdigest(),),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["request_id"] == message_id, rows
assert rows[0]["status"] == "success", rows
assert rows[0]["prompt_tokens"] == INPUT_TOKENS, rows
assert rows[0]["completion_tokens"] == FULL_OUTPUT_TOKENS, rows
assert float(str(rows[0]["spend"])) == pytest.approx(
INPUT_TOKENS * INPUT_RATE + FULL_OUTPUT_TOKENS * OUTPUT_RATE
), rows
assert len(wire.drain()) == 1

View file

@ -0,0 +1,50 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
@pytest.mark.covers("spend.failed_dispatch.failure_row_records_estimated_input_tokens")
def test_provider_500_after_dispatch_records_estimated_prompt_tokens_on_failure_row(gateway: Gateway) -> None:
prompt: Final = "failed dispatch accounting " + uuid.uuid4().hex
system: Final = "You are a terse accounting assistant"
def provider(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/chat/completions"
body: Final = json.loads(request.body)
assert body["model"] == "gpt-4o-mini"
assert body["messages"] == [{"role": "system", "content": system}, {"role": "user", "content": prompt}]
return Reply(
status=500,
body=b'{"error":{"message":"synthetic provider outage","type":"server_error","code":"500"}}',
)
with wire_server(provider) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
key: Final = scenario.key(models=[model])
failed: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "system", "content": system}, {"role": "user", "content": prompt}]},
key=key,
)
assert failed.status_code == 500 and "synthetic provider outage" in failed.text, failed.text
call_id: Final = failed.headers["x-litellm-call-id"]
assert len(wire.drain()) == 1
rows: Final = eventually(
lambda: read_rows(
'SELECT status, spend, prompt_tokens, completion_tokens, total_tokens FROM "LiteLLM_SpendLogs" '
"WHERE request_id=%s",
(call_id,),
),
lambda values: len(values) == 1,
seconds=70,
)
row: Final = rows[0]
assert row["status"] == "failure" and float(row["spend"]) == 0 and row["completion_tokens"] == 0, row
assert row["prompt_tokens"] > 0, f"failure row lost the dispatched input tokens: {row}"
assert row["total_tokens"] == row["prompt_tokens"], row

View file

@ -0,0 +1,74 @@
import json
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
ROUTER_DEPLOYMENT: Final = "router-deploy"
SELECTED_MODEL: Final = "grok-4-1-fast-reasoning"
SELECTED_MODEL_WITH_PROVIDER: Final = f"azure_ai/{SELECTED_MODEL}"
@pytest.mark.covers("spend.model_router.selected_model_is_returned_and_persisted_for_plain_alias")
def test_model_router_alias_without_router_in_name_keeps_selected_model_in_response_and_spend_log(
gateway: Gateway,
) -> None:
prompt: Final = uuid.uuid4().hex
def provider(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/chat/completions", request.target
assert json.loads(request.body) == {
"model": ROUTER_DEPLOYMENT,
"messages": [{"role": "user", "content": prompt}],
"stream": False,
}, request.body
return Reply(
body=json.dumps(
{
"id": "chatcmpl-" + uuid.uuid4().hex,
"object": "chat.completion",
"created": 1,
"model": SELECTED_MODEL,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "routed answer"},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 20, "completion_tokens": 20, "total_tokens": 40},
}
).encode()
)
with wire_server(provider) as wire, gateway.scenario() as scenario:
alias: Final = scenario.model(
model=f"azure_ai/model_router/{ROUTER_DEPLOYMENT}", api_base=wire.url, num_retries=0
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": alias, "messages": [{"role": "user", "content": prompt}]},
)
assert response.status_code == 200, response.text
body: Final = response.json()
assert body["model"] == SELECTED_MODEL_WITH_PROVIDER, response.text
assert body["choices"][0]["message"]["content"] == "routed answer", response.text
assert len(wire.drain()) == 1
rows: Final = eventually(
lambda: read_rows(
'SELECT model, model_group, status FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(body["id"],),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows == [{"model": SELECTED_MODEL_WITH_PROVIDER, "model_group": alias, "status": "success"}]
logs: Final = gateway.request("GET", "/spend/logs", params={"request_id": body["id"]})
assert logs.status_code == 200, logs.text
assert [(row["model"], row["model_group"]) for row in logs.json()] == [(SELECTED_MODEL_WITH_PROVIDER, alias)], (
logs.text
)

View file

@ -0,0 +1,80 @@
import os
import uuid
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, eventually, string_value
from integration._support.database import read_rows
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
def _cli_session_token(user_id: str, team_id: str) -> str:
cli_user: Final = LiteLLM_UserTable(user_id=user_id, user_role="internal_user", teams=[team_id], models=[])
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team_id, team_alias="cli-team")
@pytest.mark.covers("quota_management.organization_budget.cli_session_token_without_org_id_charges_team_organization")
def test_cli_session_token_without_org_id_charges_and_caps_the_team_organization(
gateway: Gateway, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("LITELLM_SALT_KEY", os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"))
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
organization: Final = gateway.post(
"/organization/new", {"organization_alias": f"integration-{uuid.uuid4().hex}", "max_budget": 0.06}
)
org_id: Final = string_value(organization["organization_id"])
scenario.cleanups.callback(
lambda: gateway.request("DELETE", "/organization/delete", {"organization_ids": [org_id]})
)
user_id: Final = scenario.user()
team_id: Final = scenario.team(
organization_id=org_id, models=[model], members_with_roles=[{"role": "user", "user_id": user_id}]
)
token: Final = _cli_session_token(user_id, team_id)
prompt: Final = f"org budget {uuid.uuid4().hex}"
upstream.get("/__observations").raise_for_status()
first: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}]},
key=token,
)
assert first.status_code == 200 and first.json()["usage"]["total_tokens"] == 40, first.text
reached_upstream: Final = upstream.get("/__observations").json()["requests"]
assert len(reached_upstream) == 1, reached_upstream
assert reached_upstream[0]["body"]["model"] == "gpt-4o-mini", reached_upstream
assert reached_upstream[0]["body"]["messages"] == [{"role": "user", "content": prompt}], reached_upstream
logged: Final = eventually(
lambda: read_rows(
'SELECT organization_id, team_id, spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(first.json()["id"],),
),
lambda values: len(values) == 1,
seconds=70,
)
assert [(row["organization_id"], row["team_id"], float(row["spend"])) for row in logged] == [
(org_id, team_id, pytest.approx(0.06))
]
charged: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_OrganizationTable" WHERE organization_id=%s', (org_id,)),
lambda values: len(values) == 1 and float(values[0]["spend"]) >= 0.06,
seconds=70,
)
assert float(charged[0]["spend"]) == pytest.approx(0.06)
assert float(gateway.get("/organization/info", {"organization_id": org_id})["spend"]) == pytest.approx(0.06)
denied: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"over org budget {uuid.uuid4().hex}"}]},
key=token,
)
assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text
assert f"Organization={org_id}" in denied.json()["error"]["message"], denied.text
assert upstream.get("/__observations").json()["requests"] == []

View file

@ -0,0 +1,110 @@
import uuid
from hashlib import sha256
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from integration._support.upstream import delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import JsonResponse
from pydantic import JsonValue
INPUT_COST_PER_TOKEN: Final = 0.000001
OUTPUT_COST_PER_TOKEN: Final = 0.001
PROMPT_TOKENS: Final = 10
CANDIDATE_TOKENS: Final = 5
COST_PER_CALL: Final = PROMPT_TOKENS * INPUT_COST_PER_TOKEN + CANDIDATE_TOKENS * OUTPUT_COST_PER_TOKEN
MAX_BUDGET: Final = 0.02
CALLS_WITHIN_BUDGET: Final = 4
def _key_spend(digest: str) -> float:
rows: Final = read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,))
assert len(rows) == 1, rows
return float(rows[0]["spend"])
def _generate_content_request(model: str) -> dict[str, JsonValue]:
return {"contents": [{"role": "user", "parts": [{"text": f"budget {model}"}]}]}
def _generate_content_response(model: str) -> JsonResponse:
return JsonResponse(
content_type="application/json",
body={
"candidates": [
{
"content": {"parts": [{"text": f"scripted answer {model}"}], "role": "model"},
"finishReason": "STOP",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": PROMPT_TOKENS,
"candidatesTokenCount": CANDIDATE_TOKENS,
"totalTokenCount": PROMPT_TOKENS + CANDIDATE_TOKENS,
},
"modelVersion": model,
},
)
def _served_call(gateway: Gateway, model: str, key: str, scenario_id: str, call: int) -> None:
digest: Final = sha256(key.encode()).hexdigest()
spend_before: Final = _key_spend(digest)
assert spend_before == pytest.approx((call - 1) * COST_PER_CALL) and spend_before < MAX_BUDGET
response: Final = gateway.request(
"POST",
f"/gemini/v1beta/models/{model}:generateContent",
_generate_content_request(model),
headers={"x-goog-api-key": key, "x-pass-x-scripted-scenario": scenario_id},
)
assert response.status_code == 200, f"call {call} with key spend {spend_before}: {response.text}"
assert response.json() == _generate_content_response(model).body, response.text
eventually(lambda: _key_spend(digest), lambda spend: spend >= call * COST_PER_CALL - 1e-9, seconds=70)
@pytest.mark.covers("spend.budget_reservation.gemini_passthrough_success_releases_reservation_from_spend_counter")
def test_repeated_gemini_passthrough_calls_stay_served_while_key_spend_is_below_max_budget(
gateway: Gateway, tmp_path: Path
) -> None:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["environment_variables"] = {
"GEMINI_API_BASE": gateway.upstream_url,
"GEMINI_API_KEY": "scripted",
}
path: Final = tmp_path / "gemini-passthrough.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
model: Final = f"gemini-passthrough-{uuid.uuid4().hex}"
created: Final = candidate.post(
"/model/new",
{
"model_name": model,
"litellm_params": {
"model": "gemini/gemini-2.5-flash",
"api_key": "scripted",
"api_base": gateway.upstream_url,
"input_cost_per_token": INPUT_COST_PER_TOKEN,
"output_cost_per_token": OUTPUT_COST_PER_TOKEN,
},
"model_info": {"id": model, "max_output_tokens": 10},
},
)
scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"]))
handle: Final = register_scenario(f"sc-{model}", _generate_content_response(model))
scenario.cleanups.callback(delete_scenario, handle)
key: Final = scenario.key(models=[model], max_budget=MAX_BUDGET)
for call in range(1, CALLS_WITHIN_BUDGET + 1):
_served_call(candidate, model, key, handle.scenario_id, call)
assert _key_spend(sha256(key.encode()).hexdigest()) == pytest.approx(CALLS_WITHIN_BUDGET * COST_PER_CALL)
denied: Final = candidate.request(
"POST",
f"/gemini/v1beta/models/{model}:generateContent",
_generate_content_request(model),
headers={"x-goog-api-key": key, "x-pass-x-scripted-scenario": handle.scenario_id},
)
assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text

View file

@ -18,3 +18,51 @@ def test_spend_calculate_rejects_unpriced_model_with_400(gateway: Gateway) -> No
assert error["type"] == "invalid_request_error", response.text
assert error["param"] == "model", response.text
assert model in string_value(error["message"]), response.text
GEMINI_LIVE_PREVIEW_MODELS: Final = (
"gemini-live-2.5-flash-preview-native-audio-09-2025",
"gemini/gemini-live-2.5-flash-preview-native-audio-09-2025",
)
@pytest.mark.parametrize("model", GEMINI_LIVE_PREVIEW_MODELS)
@pytest.mark.covers("quota_management.spend_tracking.spend_calculate.live_preview_cached_tokens_cost_fresh_rate")
def test_live_preview_entry_charges_cached_tokens_at_the_fresh_rate(gateway: Gateway, model: str) -> None:
def cost_with_cached_tokens(cached_tokens: int) -> float:
response: Final = gateway.request(
"POST",
"/spend/calculate",
{
"completion_response": {
"id": "chatcmpl-live-preview",
"object": "chat.completion",
"created": 1677652288,
"model": model,
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "live preview answer"},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 101_000,
"completion_tokens": 0,
"total_tokens": 101_000,
"prompt_tokens_details": {"cached_tokens": cached_tokens},
},
}
},
)
assert response.status_code == 200, response.text
cost: Final = object_value(JSON_OBJECT.validate_json(response.text))["cost"]
assert isinstance(cost, int | float)
return float(cost)
cached_cost: Final = cost_with_cached_tokens(100_000)
fresh_cost: Final = cost_with_cached_tokens(0)
assert fresh_cost > 0, fresh_cost
assert cached_cost == pytest.approx(fresh_cost), (
f"the entry publishes no cached rate, so 100k cached tokens must bill like fresh ones: {cached_cost} vs {fresh_cost}"
)

View file

@ -0,0 +1,81 @@
import uuid
from datetime import datetime, timedelta, timezone
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
@pytest.mark.covers("quota_management.spend_tracking.team_daily_activity_aggregated_reports_whole_range_team_spend")
def test_aggregated_team_activity_reports_the_whole_range_team_spend_in_one_page(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
keys: Final = tuple(scenario.key(team_id=team, models=[model]) for _ in range(2))
digests: Final = tuple(sha256(key.encode()).hexdigest() for key in keys)
for key in keys:
for _ in range(2):
reply: Final = gateway.chat(model, key=key, text=f"team activity {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
logged: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE team_id=%s', (team,)),
lambda values: len(values) == 4,
seconds=70,
)
assert sum(float(row["spend"]) for row in logged) == pytest.approx(0.24)
daily: Final = eventually(
lambda: read_rows(
'SELECT api_key, spend, successful_requests FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)
),
lambda values: sum(float(row["spend"]) for row in values) >= 0.24 - 1e-9,
seconds=70,
)
assert sorted(row["api_key"] for row in daily) == sorted(digests), daily
assert all(float(row["spend"]) == pytest.approx(0.12) and row["successful_requests"] == 2 for row in daily)
today: Final = datetime.now(timezone.utc)
response: Final = gateway.request(
"GET",
"/team/daily/activity/aggregated",
params={
"team_ids": team,
"start_date": (today - timedelta(days=1)).strftime("%Y-%m-%d"),
"end_date": (today + timedelta(days=1)).strftime("%Y-%m-%d"),
"timezone": "0",
},
)
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
metadata: Final = object_value(body["metadata"])
assert (
metadata["total_spend"],
metadata["total_prompt_tokens"],
metadata["total_completion_tokens"],
metadata["total_tokens"],
metadata["total_api_requests"],
metadata["total_successful_requests"],
metadata["total_failed_requests"],
metadata["page"],
metadata["total_pages"],
metadata["has_more"],
) == (pytest.approx(0.24), 80, 80, 160, 4, 4, 0, 1, 1, False), response.text
results: Final = body["results"]
assert isinstance(results, list) and len(results) == 1, response.text
day: Final = object_value(results[0])
assert object_value(day["metrics"])["spend"] == pytest.approx(0.24), response.text
entities: Final = object_value(object_value(day["breakdown"])["entities"])
assert set(entities) == {team}, response.text
team_bucket: Final = object_value(entities[team])
team_metrics: Final = object_value(team_bucket["metrics"])
assert (team_metrics["spend"], team_metrics["api_requests"], team_metrics["successful_requests"]) == (
pytest.approx(0.24),
4,
4,
), response.text
per_key: Final = object_value(team_bucket["api_key_breakdown"])
assert set(per_key) == set(digests), response.text
assert tuple(object_value(object_value(per_key[digest])["metrics"])["spend"] for digest in digests) == (
pytest.approx(0.12),
pytest.approx(0.12),
), response.text

View file

@ -0,0 +1,54 @@
import uuid
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
@pytest.mark.covers("spend.team_member.member_without_budget_gets_membership_row_and_spend")
def test_member_added_without_any_budget_is_charged_on_its_membership_row(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
user: Final = scenario.user()
added: Final = gateway.request(
"POST", "/team/member_add", {"team_id": team, "member": {"user_id": user, "role": "user"}}
)
assert added.status_code == 200, added.text
memberships: Final = added.json()["updated_team_memberships"]
assert [
{"user_id": row["user_id"], "team_id": row["team_id"], "budget_id": row["budget_id"], "spend": row["spend"]}
for row in memberships
] == [{"user_id": user, "team_id": team, "budget_id": None, "spend": 0}], added.text
assert read_rows(
'SELECT budget_id, spend, total_spend FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=%s',
(team, user),
) == [{"budget_id": None, "spend": 0.0, "total_spend": 0.0}]
key: Final = scenario.key(team_id=team, user_id=user, models=[model])
assert gateway.chat(model, key=key, text=f"member spend {uuid.uuid4().hex}")["usage"]["total_tokens"] == 40
charged: Final = eventually(
lambda: read_rows(
'SELECT spend, total_spend FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=%s',
(team, user),
),
lambda values: len(values) == 1 and float(values[0]["spend"]) >= 0.06,
seconds=70,
)
assert float(charged[0]["spend"]) == pytest.approx(0.06)
assert float(charged[0]["total_spend"]) == pytest.approx(0.06)
team_rows: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id=%s', (team,)),
lambda values: len(values) == 1 and float(values[0]["spend"]) >= 0.06,
seconds=70,
)
assert float(team_rows[0]["spend"]) == pytest.approx(0.06)
info: Final = gateway.get("/team/info", {"team_id": team})
listed: Final = info["team_memberships"]
assert isinstance(listed, list)
exposed: Final = [
(object_value(row)["user_id"], object_value(row)["spend"])
for row in listed
if object_value(row)["user_id"] == user
]
assert len(exposed) == 1 and exposed[0][1] == pytest.approx(0.06), info

View file

@ -2,25 +2,47 @@ import asyncio
import json
import threading
import uuid
from pathlib import Path
from typing import Final
import pytest
from hypothesis import Phase, example, given, settings, strategies as st
from openai import OpenAI
import yaml
from hypothesis import Phase, example, given, settings
from hypothesis import strategies as st
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from integration._support.wire import Reply, wire_server
from openai import OpenAI
def frame(identity: str, delta: dict, *, finish: str | None = None) -> bytes:
value: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini", "choices": [{"index": 0, "delta": delta, "finish_reason": finish}]}
value: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "delta": delta, "finish_reason": finish}],
}
return b"data: " + json.dumps(value, ensure_ascii=False).encode() + b"\n\n"
def text_stream(identity: str) -> tuple[bytes, ...]:
usage: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini", "choices": [], "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}}
return (frame(identity, {"role": "assistant", "content": "Hello "}), frame(identity, {"content": "雪 café"}), frame(identity, {}, finish="stop"), b"data: " + json.dumps(usage).encode() + b"\n\n", b"data: [DONE]\n\n")
usage: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
return (
frame(identity, {"role": "assistant", "content": "Hello "}),
frame(identity, {"content": "雪 café"}),
frame(identity, {}, finish="stop"),
b"data: " + json.dumps(usage).encode() + b"\n\n",
b"data: [DONE]\n\n",
)
@pytest.mark.covers("other.streaming.byte_partitions.preserve_text_identity_and_usage")
@ -37,14 +59,27 @@ def test_generated_tcp_partitions_preserve_unicode_text_identity_and_final_usage
boundaries: Final = (0, *sorted(cuts), len(body))
pieces: Final = tuple(body[left:right] for left, right in zip(boundaries, boundaries[1:]))
with wire_server(lambda request: Reply(content_type="text/event-stream", chunks=pieces)) as wire:
stream: Final = litellm.completion(model="openai/gpt-4o-mini", api_base=wire.url + "/v1", api_key="synthetic-stream-key", messages=[{"role": "user", "content": "partition control"}], stream=True, stream_options={"include_usage": True}, timeout=5, num_retries=0)
stream: Final = litellm.completion(
model="openai/gpt-4o-mini",
api_base=wire.url + "/v1",
api_key="synthetic-stream-key",
messages=[{"role": "user", "content": "partition control"}],
stream=True,
stream_options={"include_usage": True},
timeout=5,
num_retries=0,
)
try:
chunks: Final = tuple(stream)
finally:
asyncio.run(stream.aclose())
assert "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) == "Hello 雪 café"
assert (
"".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) == "Hello 雪 café"
)
assert {chunk.id for chunk in chunks} == {"stream-partition-control"}
assert [choice.finish_reason for chunk in chunks for choice in chunk.choices if choice.finish_reason] == ["stop"]
assert [choice.finish_reason for chunk in chunks for choice in chunk.choices if choice.finish_reason] == [
"stop"
]
usages: Final = tuple(chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None)
assert len(usages) == 1
assert usages[0].prompt_tokens == 11 and usages[0].completion_tokens == 4
@ -59,25 +94,63 @@ def test_fragmented_tool_names_and_arguments_keep_each_call_identity() -> None:
identity: Final = "stream-tools-control"
deltas: Final = (
{"role": "assistant", "tool_calls": [{"index": 0, "id": "call-add", "type": "function", "function": {"name": "ad", "arguments": ""}}, {"index": 1, "id": "call-multiply", "type": "function", "function": {"name": "multi", "arguments": ""}}]},
{"tool_calls": [{"index": 1, "function": {"name": "ply", "arguments": '{"x":3,'}}, {"index": 0, "function": {"arguments": '{"x":1,'}}]},
{"tool_calls": [{"index": 0, "function": {"name": "d", "arguments": '"y":2}'}}, {"index": 1, "function": {"arguments": '"y":4}'}}]},
{
"role": "assistant",
"tool_calls": [
{"index": 0, "id": "call-add", "type": "function", "function": {"name": "ad", "arguments": ""}},
{"index": 1, "id": "call-multiply", "type": "function", "function": {"name": "multi", "arguments": ""}},
],
},
{
"tool_calls": [
{"index": 1, "function": {"name": "ply", "arguments": '{"x":3,'}},
{"index": 0, "function": {"arguments": '{"x":1,'}},
]
},
{
"tool_calls": [
{"index": 0, "function": {"name": "d", "arguments": '"y":2}'}},
{"index": 1, "function": {"arguments": '"y":4}'}},
]
},
)
frames: Final = (
*tuple(frame(identity, delta) for delta in deltas),
frame(identity, {}, finish="tool_calls"),
b"data: [DONE]\n\n",
)
frames: Final = (*tuple(frame(identity, delta) for delta in deltas), frame(identity, {}, finish="tool_calls"), b"data: [DONE]\n\n")
with wire_server(lambda request: Reply(content_type="text/event-stream", chunks=frames)) as wire:
stream: Final = litellm.completion(model="openai/gpt-4o-mini", api_base=wire.url + "/v1", api_key="synthetic-stream-key", messages=[{"role": "user", "content": "tool control"}], stream=True, timeout=5, num_retries=0)
stream: Final = litellm.completion(
model="openai/gpt-4o-mini",
api_base=wire.url + "/v1",
api_key="synthetic-stream-key",
messages=[{"role": "user", "content": "tool control"}],
stream=True,
timeout=5,
num_retries=0,
)
try:
chunks: Final = tuple(stream)
finally:
asyncio.run(stream.aclose())
events: Final = tuple((choice.index, tool) for chunk in chunks for choice in chunk.choices for tool in (choice.delta.tool_calls or ()))
for index, name, call_id, arguments in ((0, "add", "call-add", {"x": 1, "y": 2}), (1, "multiply", "call-multiply", {"x": 3, "y": 4})):
events: Final = tuple(
(choice.index, tool)
for chunk in chunks
for choice in chunk.choices
for tool in (choice.delta.tool_calls or ())
)
for index, name, call_id, arguments in (
(0, "add", "call-add", {"x": 1, "y": 2}),
(1, "multiply", "call-multiply", {"x": 3, "y": 4}),
):
selected: Final = tuple(tool for choice, tool in events if (choice, tool.index) == (0, index))
assert "".join(tool.id or "" for tool in selected) == call_id
assert "".join(tool.function.name or "" for tool in selected) == name
assert json.loads("".join(tool.function.arguments or "" for tool in selected)) == arguments
assert {tool.index for _, tool in events} == {0, 1}
assert [choice.finish_reason for chunk in chunks for choice in chunk.choices if choice.finish_reason] == ["tool_calls"]
assert [choice.finish_reason for chunk in chunks for choice in chunk.choices if choice.finish_reason] == [
"tool_calls"
]
assert len(wire.drain()) == 1
@ -86,13 +159,27 @@ def test_proxy_stream_usage_visibility_keeps_exact_persisted_charge(gateway: Gat
with gateway.scenario() as scenario:
for include in (None, False, True):
identity: Final = "stream-usage-" + uuid.uuid4().hex
with wire_server(lambda request, identity=identity: Reply(content_type="text/event-stream", chunks=text_stream(identity))) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", input_cost_per_token=0.001, output_cost_per_token=0.002)
with OpenAI(api_key=gateway.key, base_url=str(gateway.client.base_url), timeout=5, max_retries=0) as client:
stream: Final = client.chat.completions.create(model=model, messages=[{"role": "user", "content": identity}], stream=True, **({} if include is None else {"stream_options": {"include_usage": include}}))
with wire_server(
lambda request, identity=identity: Reply(content_type="text/event-stream", chunks=text_stream(identity))
) as wire:
model: Final = scenario.model(
api_base=wire.url + "/v1", input_cost_per_token=0.001, output_cost_per_token=0.002
)
with OpenAI(
api_key=gateway.key, base_url=str(gateway.client.base_url), timeout=5, max_retries=0
) as client:
stream: Final = client.chat.completions.create(
model=model,
messages=[{"role": "user", "content": identity}],
stream=True,
**({} if include is None else {"stream_options": {"include_usage": include}}),
)
with stream:
chunks: Final = tuple(stream)
assert "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) == "Hello 雪 café"
assert (
"".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices)
== "Hello 雪 café"
)
assert {chunk.id for chunk in chunks} == {identity}
usages: Final = tuple(chunk.usage for chunk in chunks if chunk.usage is not None)
assert len(usages) == (1 if include else 0)
@ -101,28 +188,369 @@ def test_proxy_stream_usage_visibility_keeps_exact_persisted_charge(gateway: Gat
requests: Final = wire.drain()
assert len(requests) == 1
assert json.loads(requests[0].body)["stream_options"]["include_usage"] is True
rows: Final = eventually(lambda identity=identity: read_rows('SELECT spend, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)), lambda values: len(values) == 1, seconds=70)
rows: Final = eventually(
lambda identity=identity: read_rows(
'SELECT spend, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(identity,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["prompt_tokens"] == 11 and rows[0]["completion_tokens"] == 4
assert float(rows[0]["spend"]) == pytest.approx(0.019)
@pytest.mark.covers("other.streaming.messages_bridge.empty_choices_usage_chunk_completes_stream")
def test_messages_stream_completes_through_trailing_empty_choices_usage_chunk(gateway: Gateway) -> None:
identity: Final = "messages-empty-choices-" + uuid.uuid4().hex
metadata: Final = (
b"data: "
+ json.dumps(
{
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [],
"prompt_filter_results": [{"prompt_index": 0, "content_filter_results": {}}],
},
ensure_ascii=False,
).encode()
+ b"\n\n"
)
frames: Final = (metadata, *text_stream(identity))
with (
wire_server(lambda request: Reply(content_type="text/event-stream", chunks=frames)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model="azure/gpt-4o-mini", api_base=wire.url + "/v1")
with gateway.client.stream(
"POST",
"/v1/messages",
json={
"model": model,
"max_tokens": 64,
"stream": True,
"messages": [{"role": "user", "content": identity}],
},
headers={"Authorization": f"Bearer {gateway.key}"},
) as response:
assert response.status_code == 200, response.read().decode()
events: Final = tuple(
json.loads(line.removeprefix("data: ")) for line in response.iter_lines() if line.startswith("data: ")
)
assert tuple(event["type"] for event in events) == (
"message_start",
"content_block_start",
"content_block_delta",
"content_block_delta",
"content_block_stop",
"message_delta",
"message_stop",
), f"observed events: {events!r}"
assert (
"".join(event["delta"]["text"] for event in events if event["type"] == "content_block_delta") == "Hello 雪 café"
)
message_delta: Final = next(event for event in events if event["type"] == "message_delta")
assert message_delta["usage"] == {"input_tokens": 11, "output_tokens": 4}
requests: Final = wire.drain()
assert len(requests) == 1
outbound: Final = json.loads(requests[0].body)
assert outbound["stream"] is True and outbound["stream_options"] == {"include_usage": True}, (
f"observed outbound body: {outbound!r}"
)
@pytest.mark.covers("other.streaming.responses_bridge.empty_choices_chunks_complete_stream")
def test_responses_stream_completes_through_empty_choices_metadata_and_usage_chunks(gateway: Gateway) -> None:
identity: Final = "responses-empty-choices-" + uuid.uuid4().hex
metadata: Final = (
b"data: "
+ json.dumps(
{
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [],
"prompt_filter_results": [{"prompt_index": 0, "content_filter_results": {}}],
},
ensure_ascii=False,
).encode()
+ b"\n\n"
)
frames: Final = (metadata, *text_stream(identity))
with (
wire_server(lambda request: Reply(content_type="text/event-stream", chunks=frames)) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model="deepseek/gpt-4o-mini", api_base=wire.url + "/v1")
with gateway.client.stream(
"POST",
"/v1/responses",
json={"model": model, "input": identity, "stream": True},
headers={"Authorization": f"Bearer {gateway.key}"},
) as response:
assert response.status_code == 200, response.read().decode()
events: Final = tuple(
json.loads(line.removeprefix("data: "))
for line in response.iter_lines()
if line.startswith("data: ") and line != "data: [DONE]"
)
assert (
"".join(event["delta"] for event in events if event["type"] == "response.output_text.delta") == "Hello 雪 café"
), f"observed events: {events!r}"
assert tuple(event["type"] for event in events if event["type"] != "response.output_text.delta") == (
"response.created",
"response.in_progress",
"response.output_item.added",
"response.content_part.added",
"response.output_text.done",
"response.content_part.done",
"response.output_item.done",
"response.completed",
), f"observed events: {events!r}"
assert events[-1]["type"] == "response.completed"
assert events[-1]["response"]["usage"] == {
"input_tokens": 11,
"output_tokens": 4,
"output_tokens_details": {"reasoning_tokens": 0, "text_tokens": 4},
"total_tokens": 15,
}
requests: Final = wire.drain()
assert len(requests) == 1
outbound: Final = json.loads(requests[0].body)
assert outbound["stream"] is True and outbound["stream_options"] == {"include_usage": True}, (
f"observed outbound body: {outbound!r}"
)
def provider_cost_object_stream(identity: str, total_cost: float) -> tuple[bytes, ...]:
cost: Final = {
"input_tokens_cost": 0.0001,
"output_tokens_cost": 0.0002,
"request_cost": 0.012,
"total_cost": total_cost,
}
usage: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "sonar",
"choices": [],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15, "cost": cost},
}
return (
frame(identity, {"role": "assistant", "content": "Hello "}),
frame(identity, {"content": "from search"}),
frame(identity, {}, finish="stop"),
b"data: " + json.dumps(usage).encode() + b"\n\n",
b"data: [DONE]\n\n",
)
def sse_data_lines(text: str) -> tuple[str, ...]:
return tuple(line.removeprefix("data: ") for line in text.splitlines() if line.startswith("data: "))
@pytest.mark.covers("other.streaming.usage.provider_cost_object_completes_stream_and_bills_total_cost")
def test_perplexity_stream_with_cost_breakdown_object_completes_and_bills_total_cost(gateway: Gateway) -> None:
identity: Final = "stream-cost-object-" + uuid.uuid4().hex
total_cost: Final = 0.0123
with (
gateway.scenario() as scenario,
wire_server(
lambda request: Reply(
content_type="text/event-stream", chunks=provider_cost_object_stream(identity, total_cost)
)
) as wire,
):
model: Final = scenario.model(model="perplexity/sonar", api_base=wire.url + "/v1")
with gateway.client.stream(
"POST",
"/v1/chat/completions",
json={
"model": model,
"messages": [{"role": "user", "content": identity}],
"stream": True,
"stream_options": {"include_usage": True},
},
headers={"Authorization": f"Bearer {gateway.key}"},
) as response:
text: Final = response.read().decode()
assert response.status_code == 200, text
lines: Final = sse_data_lines(text)
assert lines[-1] == "[DONE]", text
events: Final = tuple(json.loads(line) for line in lines[:-1])
assert [event for event in events if "error" in event] == [], text
assert (
"".join(choice["delta"].get("content") or "" for event in events for choice in event["choices"])
== "Hello from search"
), text
assert [
choice.get("finish_reason")
for event in events
for choice in event["choices"]
if choice.get("finish_reason")
] == ["stop"], text
usages: Final = tuple(event["usage"] for event in events if event.get("usage") is not None)
assert len(usages) == 1, text
assert (usages[0]["prompt_tokens"], usages[0]["completion_tokens"], usages[0]["total_tokens"]) == (11, 4, 15), (
text
)
requests: Final = wire.drain()
assert len(requests) == 1
outbound: Final = json.loads(requests[0].body)
assert outbound["model"] == "sonar" and outbound["stream"] is True, outbound
assert outbound["messages"] == [{"role": "user", "content": identity}], outbound
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(identity,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert (rows[0]["prompt_tokens"], rows[0]["completion_tokens"]) == (11, 4)
assert float(rows[0]["spend"]) == pytest.approx(total_cost)
@pytest.mark.covers(
"other.streaming.fallback.empty_leading_chunk_then_disconnect_streams_fallback_with_usage_and_spend"
)
def test_primary_stream_with_empty_first_chunk_then_disconnect_falls_back_and_bills_the_fallback(
gateway: Gateway,
tmp_path: Path,
) -> None:
identity: Final = "stream-empty-fallback-" + uuid.uuid4().hex
empty_first: Final = (
b"data: "
+ json.dumps(
{
"id": identity + "-primary",
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [],
"usage": {"prompt_tokens": 11, "completion_tokens": 0, "total_tokens": 11},
}
).encode()
+ b"\n\n"
)
with (
wire_server(
lambda request: Reply(
content_type="text/event-stream",
chunks=(empty_first, b":" + b"x" * 4_000_000 + b"\n\n", empty_first),
abort_after=2,
)
) as primary,
wire_server(
lambda request: Reply(content_type="text/event-stream", chunks=text_stream(identity))
) as fallback,
):
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
{
"model_name": name,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "synthetic-fallback-key",
"api_base": server.url + "/v1",
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
},
}
for name, server in (("primary", primary), ("fallback", fallback))
]
config["router_settings"] = {
"num_retries": 0,
"disable_cooldowns": True,
"fallbacks": [{"primary": ["fallback"]}],
}
path: Final = tmp_path / "fallbacks.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate:
body: Final = {
"model": "primary",
"messages": [{"role": "user", "content": identity}],
"stream": True,
"stream_options": {"include_usage": True},
}
with candidate.client.stream(
"POST", "/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {candidate.key}"}
) as response:
lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data:"))
assert response.status_code == 200, lines
assert lines[-1] == "data: [DONE]", lines
events: Final = tuple(json.loads(line.removeprefix("data:")) for line in lines[:-1])
assert all("error" not in event for event in events), lines
assert (
"".join(choice["delta"].get("content") or "" for event in events for choice in event["choices"])
== "Hello 雪 café"
), lines
usages: Final = tuple(event["usage"] for event in events if event.get("usage") is not None)
assert (usages[-1]["prompt_tokens"], usages[-1]["completion_tokens"]) == (11, 4), lines
assert tuple(
json.loads(request.body)["messages"]
for request in primary.drain()
if request.target.endswith("/chat/completions")
) == (body["messages"],)
assert tuple(
json.loads(request.body)["messages"]
for request in fallback.drain()
if request.target.endswith("/chat/completions")
) == (body["messages"],)
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, prompt_tokens, completion_tokens, status FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(identity,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert (rows[0]["prompt_tokens"], rows[0]["completion_tokens"], rows[0]["status"]) == (11, 4, "success"), (
rows
)
assert float(rows[0]["spend"]) == pytest.approx(0.019), rows
@pytest.mark.covers("other.streaming.failure.truncated_transport_raises_and_control_recovers")
def test_truncated_http_stream_is_an_error_and_next_stream_succeeds() -> None:
import litellm
for truncated in (True, False):
with wire_server(lambda request, truncated=truncated: Reply(content_type="text/event-stream", chunks=text_stream("stream-truncated"), abort_after=1 if truncated else None)) as wire:
stream: Final = litellm.completion(model="openai/gpt-4o-mini", api_base=wire.url + "/v1", api_key="synthetic-stream-key", messages=[{"role": "user", "content": "truncation control"}], stream=True, timeout=5, num_retries=0)
with wire_server(
lambda request, truncated=truncated: Reply(
content_type="text/event-stream",
chunks=text_stream("stream-truncated"),
abort_after=1 if truncated else None,
)
) as wire:
stream: Final = litellm.completion(
model="openai/gpt-4o-mini",
api_base=wire.url + "/v1",
api_key="synthetic-stream-key",
messages=[{"role": "user", "content": "truncation control"}],
stream=True,
timeout=5,
num_retries=0,
)
try:
if truncated:
with pytest.raises(litellm.exceptions.MidStreamFallbackError, match="incomplete chunked read") as failure:
with pytest.raises(
litellm.exceptions.MidStreamFallbackError, match="incomplete chunked read"
) as failure:
tuple(stream)
assert isinstance(failure.value.original_exception, litellm.APIConnectionError)
assert failure.value.generated_content == "Hello "
assert failure.value.is_pre_first_chunk is False
else:
chunks: Final = tuple(stream)
assert "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) == "Hello 雪 café"
assert (
"".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices)
== "Hello 雪 café"
)
assert any(choice.finish_reason == "stop" for chunk in chunks for choice in chunk.choices)
finally:
asyncio.run(stream.aclose())
@ -134,9 +562,23 @@ def test_client_cancellation_releases_the_actual_provider_connection() -> None:
import litellm
gate: Final = threading.Event()
frames: Final = (frame("stream-cancel", {"role": "assistant", "content": "first"}), b":" + b"x" * 4_000_000 + b"\n\n", b"data: [DONE]\n\n")
with wire_server(lambda request: Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate)) as wire:
stream: Final = litellm.completion(model="openai/gpt-4o-mini", api_base=wire.url + "/v1", api_key="synthetic-stream-key", messages=[{"role": "user", "content": "cancellation control"}], stream=True, timeout=5, num_retries=0)
frames: Final = (
frame("stream-cancel", {"role": "assistant", "content": "first"}),
b":" + b"x" * 4_000_000 + b"\n\n",
b"data: [DONE]\n\n",
)
with wire_server(
lambda request: Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate)
) as wire:
stream: Final = litellm.completion(
model="openai/gpt-4o-mini",
api_base=wire.url + "/v1",
api_key="synthetic-stream-key",
messages=[{"role": "user", "content": "cancellation control"}],
stream=True,
timeout=5,
num_retries=0,
)
try:
first: Final = next(stream)
assert first.choices[0].delta.content == "first"

View file

@ -0,0 +1,86 @@
import json
import uuid
from pathlib import Path
from typing import Final
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
def frame(identity: str, delta: dict[str, str], *, finish: str | None = None) -> bytes:
event: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "delta": delta, "finish_reason": finish}],
}
return b"data: " + json.dumps(event).encode() + b"\n\n"
@pytest.mark.covers("streaming.max_parallel_requests.slot_released_when_stream_logging_callback_fails")
def test_failing_stream_logging_callback_does_not_leak_max_parallel_requests_slot(
gateway: Gateway, tmp_path: Path
) -> None:
identity: Final = "stream-slot-" + uuid.uuid4().hex
prompt: Final = "slot release control " + identity
def analyzer(request: Request) -> Reply:
assert request.target == "/analyze"
assert json.loads(request.body)["text"] == prompt
return Reply(status=500, body=json.dumps({"error": "synthetic analyzer outage"}).encode())
def provider(request: Request) -> Reply:
assert request.target == "/v1/chat/completions"
body: Final = json.loads(request.body)
assert body["messages"] == [{"role": "user", "content": prompt}]
assert body["stream"] is True
return Reply(
content_type="text/event-stream",
chunks=(
frame(identity, {"role": "assistant", "content": "Hello"}),
frame(identity, {"content": " slot"}),
frame(identity, {}, finish="stop"),
b"data: [DONE]\n\n",
),
)
with wire_server(analyzer) as policy, wire_server(provider) as upstream:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": identity,
"litellm_params": {
"guardrail": "presidio",
"mode": "logging_only",
"default_on": True,
"presidio_filter_scope": "input",
"pii_entities_config": {"EMAIL_ADDRESS": "MASK"},
"presidio_analyzer_api_base": policy.url + "/",
"presidio_anonymizer_api_base": policy.url + "/",
},
}
]
path: Final = tmp_path / "failing_logging_guardrail.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
model: Final = scenario.model(api_base=upstream.url + "/v1")
key: Final = scenario.key(max_parallel_requests=1)
body: Final = {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True}
first: Final = candidate.request("POST", "/v1/chat/completions", body, key=key)
assert first.status_code == 200, first.text
assert first.text.endswith("data: [DONE]\n\n"), first.text
assert len(upstream.drain()) == 1
eventually(lambda: policy.received.qsize(), lambda count: count >= 1)
assert {scan.target for scan in policy.drain()} == {"/analyze"}
second: Final = eventually(
lambda: candidate.request("POST", "/v1/chat/completions", body, key=key),
lambda response: response.status_code == 200,
seconds=20,
return_last_on_timeout=True,
)
assert second.status_code == 200, second.text
assert second.text.endswith("data: [DONE]\n\n"), second.text

View file

@ -0,0 +1,62 @@
import json
import threading
import uuid
from collections.abc import Callable, Iterable, Iterator
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.wire import Reply, Request, wire_server
from integration.streaming.test_stream_contracts import text_stream
KEEPALIVE_SECONDS: Final = 1
def _reply_after_first_ping(identity: str, first_ping_seen: threading.Event) -> Callable[[Request], Reply]:
def respond(_request: Request) -> Reply:
first_ping_seen.wait(timeout=10)
return Reply(content_type="text/event-stream", chunks=text_stream(identity))
return respond
def _frames_setting(first_ping_seen: threading.Event, lines: Iterable[str]) -> Iterator[str]:
for line in lines:
if line == ": ping":
first_ping_seen.set()
yield line
@pytest.mark.covers("streaming.keepalive.sse_pings_fill_silent_time_to_first_token")
def test_stream_emits_sse_ping_comments_before_the_first_data_frame_while_upstream_is_silent(
gateway: Gateway,
) -> None:
identity: Final = "stream-ttft-keepalive-" + uuid.uuid4().hex
first_ping_seen: Final = threading.Event()
with gateway.scenario() as scenario:
with wire_server(_reply_after_first_ping(identity, first_ping_seen)) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", keepalive_seconds=KEEPALIVE_SECONDS)
with gateway.client.stream(
"POST",
"/v1/chat/completions",
json={"model": model, "messages": [{"role": "user", "content": identity}], "stream": True},
headers={"Authorization": f"Bearer {gateway.key}"},
) as response:
assert response.status_code == 200, response.read().decode()
frames: Final = tuple(
_frames_setting(first_ping_seen, (line for line in response.iter_lines() if line))
)
first_data: Final = next(index for index, line in enumerate(frames) if line.startswith("data:"))
assert first_data >= 1, f"No keepalive reached the client before the first data frame: {frames}"
assert frames[:first_data] == (": ping",) * first_data, frames
assert frames[-1] == "data: [DONE]", frames
deltas: Final = tuple(json.loads(line.removeprefix("data: ")) for line in frames[first_data:-1])
assert (
"".join(choice["delta"].get("content", "") for chunk in deltas for choice in chunk["choices"])
== "Hello 雪 café"
), frames
requests: Final = wire.drain()
assert len(requests) == 1
outbound: Final = json.loads(requests[0].body)
assert outbound["model"] == "gpt-4o-mini" and outbound["stream"] is True, outbound
assert "keepalive_seconds" not in outbound, outbound

View file

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

View file

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

View file

@ -1,17 +1,20 @@
import asyncio
import hashlib
import json
from typing import Iterable, List, Optional, Tuple
import time
from collections.abc import Iterable
from unittest.mock import patch
import pytest
from redis.asyncio import Redis
import litellm.proxy.common_utils.auth_cache_invalidation_pubsub as pubsub_module
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
AUTH_CACHE_INVALIDATION_CHANNEL,
AuthCacheInvalidationSubscriber,
evict_and_broadcast,
publish_auth_cache_invalidation,
)
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
@ -19,13 +22,29 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
class _RecordingRedisClient(Redis):
def __init__(self) -> None:
self.published: List[Tuple[str, str]] = []
self.published: list[tuple[str, str]] = []
async def publish(self, channel: str, message: str) -> int:
self.published.append((channel, message))
return 1
class _WedgedPublishRedisClient(Redis):
def __init__(self) -> None:
self.attempted: list[str] = []
self.in_flight = 0
self.max_in_flight = 0
self.release = asyncio.Event()
async def publish(self, channel: str, message: str) -> int:
self.in_flight += 1
self.max_in_flight = max(self.max_in_flight, self.in_flight)
self.attempted.append(message)
await self.release.wait()
self.in_flight -= 1
return 1
class _FailingPublishRedisClient(Redis):
def __init__(self) -> None:
pass
@ -36,16 +55,16 @@ class _FailingPublishRedisClient(Redis):
class _QueuePubSub:
def __init__(self, initial_messages: Iterable[object] = ()) -> None:
self.queue: "asyncio.Queue[object]" = asyncio.Queue()
self.queue: asyncio.Queue[object] = asyncio.Queue()
for message in initial_messages:
self.queue.put_nowait(message)
self.subscribed_channels: List[str] = []
self.subscribed_channels: list[str] = []
self.closed = False
async def subscribe(self, *channels: str) -> None:
self.subscribed_channels.extend(channels)
async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> Optional[object]:
async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> object | None:
try:
return await asyncio.wait_for(self.queue.get(), timeout)
except asyncio.TimeoutError:
@ -64,7 +83,7 @@ class _ScriptedPubSubRedisClient(Redis):
class _FakeRedisCache:
def __init__(self, client: object, namespace: Optional[str] = None) -> None:
def __init__(self, client: object, namespace: str | None = None) -> None:
self._client = client
self.namespace = namespace
@ -222,3 +241,49 @@ async def test_subscriber_ignores_malformed_messages() -> None:
subscriber._apply_message(None)
assert cache.in_memory_cache.get_cache("project_id:p-1") is not None
@pytest.mark.asyncio
async def test_evict_and_broadcast_evicts_locally_and_returns_while_redis_publish_never_answers() -> None:
cache = UserApiKeyCache()
cache.set_cache("user-wedged", UserAPIKeyAuth(user_id="user-wedged"), model_type=UserAPIKeyAuth)
client = _WedgedPublishRedisClient()
with patch(
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
return_value=_FakeRedisCache(client=client),
):
started = time.monotonic()
await evict_and_broadcast(cache_keys=("user-wedged",), user_api_key_cache=cache)
elapsed = time.monotonic() - started
assert elapsed < 0.1, f"handler waited {elapsed:.3f}s on a publish that never answers"
assert cache.get_cache("user-wedged", model_type=UserAPIKeyAuth) is None
assert client.attempted == [json.dumps({"cache_key": "user-wedged"})], "publish was not handed to redis"
client.release.set()
await asyncio.sleep(0)
@pytest.mark.asyncio
async def test_publish_holds_at_most_sixteen_redis_connections_while_redis_is_wedged(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(pubsub_module, "_in_flight_publishes", asyncio.Semaphore(16))
monkeypatch.setattr(pubsub_module, "_pending_publishes", set())
client = _WedgedPublishRedisClient()
with patch(
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
return_value=_FakeRedisCache(client=client),
):
for i in range(64):
await publish_auth_cache_invalidation(cache_key=f"user-{i}")
await asyncio.sleep(0)
await asyncio.sleep(0)
assert client.max_in_flight == 16, f"publish tasks held {client.max_in_flight} redis connections at once"
assert len(client.attempted) == 16, "waiters called publish before a semaphore slot freed"
client.release.set()
await asyncio.gather(*pubsub_module._pending_publishes) # pyright: ignore[reportPrivateUsage] # drain module-level tasks
assert len(client.attempted) == 64

View file

@ -6,14 +6,23 @@ Tests:
- Scope matching via attachments (teams, keys, models)
"""
import logging
from typing import Final
import pytest
from hypothesis import given, settings
from hypothesis import strategies as st
import litellm.proxy.policy_engine.attachment_registry as attachment_registry_module
import litellm.proxy.policy_engine.policy_registry as policy_registry_module
from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
from litellm.proxy.policy_engine.policy_registry import PolicyRegistry
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
from litellm.types.proxy.policy_engine import (
Policy,
PolicyCondition,
PolicyGuardrails,
PolicyMatchContext,
PolicyScope,
)
@ -221,6 +230,34 @@ def _global_registries(monkeypatch):
return policies
def _inherited_registries(monkeypatch, parent_condition=None):
policies = PolicyRegistry()
policies.load_policies(
{
"parent": {
"guardrails": {"add": ["y"]},
**({"condition": parent_condition} if parent_condition else {}),
},
"child": {
"inherit": "parent",
"guardrails": {"add": ["x"]},
"condition": {"model": "claude.*"},
},
"fallback": {"guardrails": {"add": ["z"]}},
}
)
attachments = AttachmentRegistry()
attachments.load_attachments(
[
{"policy": "child", "scope": "*"},
{"policy": "fallback", "scope": "*", "default": True},
]
)
monkeypatch.setattr(policy_registry_module, "get_policy_registry", lambda: policies)
monkeypatch.setattr(attachment_registry_module, "get_attachment_registry", lambda: attachments)
return policies
class TestGetMatchingPoliciesFallback:
def test_condition_failing_opt_in_falls_back_to_default(self, monkeypatch):
_global_registries(monkeypatch)
@ -244,3 +281,154 @@ class TestGetMatchingPoliciesFallback:
PolicyMatcher.get_matching_policies(context=context)
assert len(calls) == 1
def test_condition_missing_child_with_unconditional_parent_still_matches(self, monkeypatch):
_inherited_registries(monkeypatch)
context = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-5.5")
assert PolicyMatcher.get_matching_policies(context=context) == ["child"]
def test_child_whose_whole_chain_misses_falls_back_to_default(self, monkeypatch):
_inherited_registries(monkeypatch, parent_condition={"model": "claude.*"})
context = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-5.5")
assert PolicyMatcher.get_matching_policies(context=context) == ["fallback"]
def test_get_policies_with_matching_conditions_keeps_missing_policy_out(self):
policies = {
"real": Policy(
guardrails=PolicyGuardrails(add=["g"]),
condition=PolicyCondition(model="claude.*"),
),
}
context = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-5.5")
assert (
PolicyMatcher.get_policies_with_matching_conditions(
policy_names=["nope"], context=context, policies=policies
)
== []
)
_MODELS: Final = ("gpt-4o", "gpt-5.5", "claude-opus-4-1")
def _policy_forest(draw: st.DrawFn) -> dict[str, Policy]: # mutable-ok: PolicyResolver takes dict[str, Policy]
names: Final = tuple(f"p{i}" for i in range(draw(st.integers(min_value=1, max_value=6))))
return { # mutable-ok: PolicyResolver takes dict[str, Policy]
name: Policy(
inherit=draw(st.sampled_from((None, *names[:i]))),
guardrails=PolicyGuardrails(add=[f"g-{name}"]), # mutable-ok: pydantic list field
condition=draw(st.sampled_from((None, *(PolicyCondition(model=m) for m in _MODELS)))),
)
for i, name in enumerate(names)
}
@st.composite
def _forest_and_request(
draw: st.DrawFn,
) -> tuple[dict[str, Policy], tuple[str, ...], PolicyMatchContext]: # mutable-ok: PolicyResolver takes dict
policies: Final = _policy_forest(draw)
attached: Final = tuple(draw(st.lists(st.sampled_from(sorted(policies)), unique=True)))
context: Final = PolicyMatchContext(team_alias="t", key_alias="k", model=draw(st.sampled_from(_MODELS)))
return policies, attached, context
def _own_condition_applies(policy: Policy, context: PolicyMatchContext) -> bool:
return policy.condition is None or policy.condition.model == context.model
def _applicable_chain(
policies: dict[str, Policy], # mutable-ok: PolicyResolver takes dict[str, Policy]
name: str,
context: PolicyMatchContext,
) -> tuple[str, ...]:
chain: Final = PolicyResolver.resolve_inheritance_chain(policy_name=name, policies=policies)
return tuple(member for member in chain if _own_condition_applies(policies[member], context))
class TestChainMatchingProperties:
@given(_forest_and_request())
@settings(max_examples=400, deadline=None)
def test_chain_matching_only_widens_to_applicable_ancestor_guardrails(
self,
case: tuple[dict[str, Policy], tuple[str, ...], PolicyMatchContext], # mutable-ok: PolicyResolver takes dict
):
policies, attached, context = case
head: Final = tuple(
PolicyMatcher.get_policies_with_matching_conditions(
policy_names=attached, context=context, policies=policies
)
)
base: Final = tuple(name for name in attached if _own_condition_applies(policies[name], context))
expected_head: Final = tuple(name for name in attached if _applicable_chain(policies, name, context))
assert head == expected_head, "a policy applies exactly when some chain member's own condition applies"
assert frozenset(base) <= frozenset(head), "head must never drop a policy base applied"
for name in head:
resolved = PolicyResolver.resolve_policy_guardrails(policy_name=name, policies=policies, context=context)
assert sorted(resolved.guardrails) == sorted(
f"g-{member}" for member in _applicable_chain(policies, name, context)
)
if name not in base:
assert f"g-{name}" not in resolved.guardrails, "a condition-missed child must not add its own guardrail"
class TestAncestorAdmissionLogging:
@staticmethod
def _chain() -> dict[str, Policy]: # mutable-ok: PolicyResolver takes dict[str, Policy]
return { # mutable-ok: PolicyResolver takes dict[str, Policy]
"parent": Policy(guardrails=PolicyGuardrails(add=["g-parent"])), # mutable-ok: pydantic list field
"child": Policy(
inherit="parent",
guardrails=PolicyGuardrails(add=["g-child"]), # mutable-ok: pydantic list field
condition=PolicyCondition(model="gpt-5.5"),
),
}
def test_logs_when_admitted_through_ancestor_only(self, caplog):
context: Final = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-4o")
with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"):
result: Final = PolicyMatcher.policy_applies(context, self._chain())("child")
records: Final = [r for r in caplog.records if "applied through ancestor" in r.getMessage()]
assert result is True
assert len(records) == 1
assert "applied through ancestor 'parent'" in records[0].getMessage()
assert "'child'" in records[0].getMessage()
def test_no_log_when_own_condition_matches(self, caplog):
context: Final = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-5.5")
with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"):
result: Final = PolicyMatcher.policy_applies(context, self._chain())("child")
assert result is True
assert not [r for r in caplog.records if "applied through ancestor" in r.getMessage()]
def test_no_log_when_no_chain_member_applies(self, caplog):
policies: Final = { # mutable-ok: PolicyResolver takes dict[str, Policy]
"parent": Policy(
guardrails=PolicyGuardrails(add=["g-parent"]), # mutable-ok: pydantic list field
condition=PolicyCondition(model="claude-opus-4-1"),
),
"child": Policy(
inherit="parent",
guardrails=PolicyGuardrails(add=["g-child"]), # mutable-ok: pydantic list field
condition=PolicyCondition(model="gpt-5.5"),
),
}
context: Final = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-4o")
with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"):
result: Final = PolicyMatcher.policy_applies(context, policies)("child")
assert result is False
assert not [r for r in caplog.records if "applied through ancestor" in r.getMessage()]
def test_condition_filter_logs_nothing(self, caplog):
context: Final = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-4o")
with caplog.at_level(logging.INFO, logger="LiteLLM Proxy"):
result: Final = PolicyMatcher.get_policies_with_matching_conditions(
policy_names=["child"], context=context, policies=self._chain()
)
assert result == ["child"]
assert not [r for r in caplog.records if "applied through ancestor" in r.getMessage()]

View file

@ -11,6 +11,8 @@ import pytest
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
from litellm.types.proxy.policy_engine import (
GuardrailPipeline,
PipelineStep,
Policy,
PolicyCondition,
PolicyGuardrails,
@ -199,3 +201,55 @@ class TestPolicyResolverWithConditions:
)
assert "pii_blocker" in resolved_gpt35.guardrails
assert "child_guardrail" not in resolved_gpt35.guardrails
def test_resolve_guardrails_for_context_with_condition_missing_child_keeps_inherited_parent(self):
"""Test a matched child whose condition misses still contributes unconditional parent guardrails."""
policies = {
"parent": Policy(
guardrails=PolicyGuardrails(add=["y"]),
),
"child": Policy(
inherit="parent",
guardrails=PolicyGuardrails(add=["x"]),
condition=PolicyCondition(model="claude.*"),
),
}
context_miss = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-5.5")
assert PolicyResolver.resolve_guardrails_for_context(
context=context_miss, policies=policies, policy_names=["child"]
) == ["y"]
context_hit = PolicyMatchContext(team_alias="t", key_alias="k", model="claude-haiku")
assert set(
PolicyResolver.resolve_guardrails_for_context(
context=context_hit, policies=policies, policy_names=["child"]
)
) == {"x", "y"}
def test_resolve_pipelines_for_context_skips_pipeline_when_own_condition_misses(self):
"""Test a matched child whose own condition misses does not run its pipeline."""
pipeline = GuardrailPipeline(mode="pre_call", steps=[PipelineStep(guardrail="child-guard")])
policies = {
"parent": Policy(
guardrails=PolicyGuardrails(add=["y"]),
),
"child": Policy(
inherit="parent",
pipeline=pipeline,
condition=PolicyCondition(model="gpt-5.5"),
),
}
context_miss = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-4o")
assert (
PolicyResolver.resolve_pipelines_for_context(
context=context_miss, policies=policies, policy_names=["child"]
)
== []
)
context_hit = PolicyMatchContext(team_alias="t", key_alias="k", model="gpt-5.5")
assert PolicyResolver.resolve_pipelines_for_context(
context=context_hit, policies=policies, policy_names=["child"]
) == [("child", pipeline)]

View file

@ -4377,6 +4377,45 @@ def test_match_and_track_policies_preserves_attachment_and_request_body_order():
assert applied_policy_names == policy_names
def test_match_and_track_policies_keeps_condition_missing_child_alongside_unconditional_sibling():
from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry
from litellm.types.proxy.policy_engine import (
Policy,
PolicyCondition,
PolicyGuardrails,
PolicyMatchContext,
)
policies = {
"baseline": Policy(guardrails=PolicyGuardrails(add=["baseline_guardrail"])),
"parent": Policy(guardrails=PolicyGuardrails(add=["pii_blocker"])),
"child": Policy(
inherit="parent",
guardrails=PolicyGuardrails(add=["child_guard"]),
condition=PolicyCondition(model="claude.*"),
),
}
attachment_registry = AttachmentRegistry()
attachment_registry.load_attachments(
[
{"policy": "baseline", "scope": "*"},
{"policy": "child", "scope": "*"},
]
)
data = {"metadata": {}}
applied_policy_names, _ = _match_and_track_policies(
data=data,
context=PolicyMatchContext(model="gpt-5.5"),
request_body_policies=[],
policies_override=policies,
attachment_registry_override=attachment_registry,
)
assert applied_policy_names == ["baseline", "child"]
assert data["metadata"]["applied_policies"] == ["baseline", "child"]
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_keeps_a_policy_added_guardrail_its_pipeline_also_steps():
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
@ -4419,6 +4458,48 @@ async def test_add_guardrails_from_policy_engine_keeps_a_policy_added_guardrail_
assert [pipeline.mode for _policy_name, pipeline in data["metadata"]["_guardrail_pipelines"]] == ["post_call"]
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_applies_inherited_parent_guardrail_when_child_condition_misses():
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
from litellm.types.proxy.policy_engine import (
Policy,
PolicyAttachment,
PolicyCondition,
PolicyGuardrails,
)
data = {"model": "gpt-5.5", "messages": [{"role": "user", "content": "Hello"}], "metadata": {}}
policy_registry = get_policy_registry()
policy_registry._policies = {
"parent": Policy(guardrails=PolicyGuardrails(add=["pii_blocker"])),
"child": Policy(
inherit="parent",
guardrails=PolicyGuardrails(add=["child_guard"]),
condition=PolicyCondition(model="claude.*"),
),
}
policy_registry._initialized = True
attachment_registry = get_attachment_registry()
attachment_registry._attachments = [PolicyAttachment(policy="child", scope="*")]
attachment_registry._initialized = True
try:
await add_guardrails_from_policy_engine(
data=data,
metadata_variable_name="metadata",
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
)
finally:
policy_registry._policies = {}
policy_registry._initialized = False
attachment_registry._attachments = []
attachment_registry._initialized = False
assert "pii_blocker" in data["metadata"]["guardrails"]
assert "child_guard" not in data["metadata"]["guardrails"]
@pytest.mark.asyncio
async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_pops_from_data():
"""

View file

@ -2189,6 +2189,7 @@ async def test_new_vector_store_persists_embedding_reference_without_credentials
mock_registry = MagicMock()
mock_registry.add_vector_store_to_registry = MagicMock()
mock_registry.is_config_vector_store.return_value = False
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
@ -2267,6 +2268,7 @@ async def test_new_vector_store_auto_resolves_from_router():
mock_registry = MagicMock()
mock_registry.add_vector_store_to_registry = MagicMock()
mock_registry.is_config_vector_store.return_value = False
with (
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
@ -3061,3 +3063,196 @@ def test_vector_store_search_rejects_caller_embedding_selection_params(blocked_k
assert response.status_code == 400, response.json()
assert blocked_key in str(response.json())
class TestConfigOwnedVectorStores:
"""Stores declared under ``vector_store_registry`` in config.yaml are owned by the config file"""
CONFIG_ID = "vs_from_config"
DB_ID = "vs_from_db"
def _registry(self):
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
registry = VectorStoreRegistry(vector_stores=[])
registry.load_vector_stores_from_config(
[
{
"vector_store_name": "config-store",
"litellm_params": {"vector_store_id": self.CONFIG_ID, "custom_llm_provider": "openai"},
}
]
)
registry.add_vector_store_to_registry(self._db_row(self.DB_ID, "db-store"))
registry.add_vector_store_to_registry(self._db_row("vs_stale", "deleted-elsewhere"))
return registry
@staticmethod
def _db_row(vector_store_id: str, vector_store_name: str) -> dict:
return {
"vector_store_id": vector_store_id,
"custom_llm_provider": "openai",
"vector_store_name": vector_store_name,
"litellm_params": {},
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
@staticmethod
def _admin() -> UserAPIKeyAuth:
return UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin")
@pytest.mark.asyncio
async def test_list_keeps_config_store_that_has_no_db_row(self):
from litellm.proxy.vector_store_endpoints.management_endpoints import list_vector_stores
registry = self._registry()
prisma = MagicMock()
prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[self._db_row(self.DB_ID, "db-store")])
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam
patch.object(litellm, "vector_store_registry", registry),
):
first = await list_vector_stores(user_api_key_dict=self._admin())
second = await list_vector_stores(user_api_key_dict=self._admin())
assert [(vs["vector_store_id"], vs["is_config"]) for vs in first["data"]] == [(self.DB_ID, False), (self.CONFIG_ID, True)]
assert second["data"] == first["data"]
assert [vs["vector_store_id"] for vs in registry.vector_stores] == [self.CONFIG_ID, self.DB_ID]
@pytest.mark.asyncio
async def test_list_keeps_config_store_and_db_row_with_same_id_does_not_overwrite_it(self):
from litellm.proxy.vector_store_endpoints.management_endpoints import list_vector_stores
registry = self._registry()
prisma = MagicMock()
prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(
return_value=[self._db_row(self.DB_ID, "db-store"), self._db_row(self.CONFIG_ID, "renamed-in-db")]
)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam
patch.object(litellm, "vector_store_registry", registry),
):
response = await list_vector_stores(user_api_key_dict=self._admin())
by_id = {vs["vector_store_id"]: vs for vs in response["data"]}
assert set(by_id) == {self.CONFIG_ID, self.DB_ID}, response
assert (by_id[self.CONFIG_ID]["vector_store_name"], by_id[self.CONFIG_ID]["is_config"]) == ("config-store", True)
assert (by_id[self.DB_ID]["vector_store_name"], by_id[self.DB_ID]["is_config"]) == ("db-store", False)
assert [vs["vector_store_id"] for vs in registry.vector_stores] == [self.CONFIG_ID, self.DB_ID]
assert registry.get_litellm_managed_vector_store_from_registry(self.CONFIG_ID)["vector_store_name"] == "config-store"
@pytest.mark.asyncio
async def test_info_reports_config_ownership(self):
from litellm.proxy.vector_store_endpoints.management_endpoints import get_vector_store_info
from litellm.types.vector_stores import VectorStoreInfoRequest
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: proxy_server global, no seam
patch.object(litellm, "vector_store_registry", self._registry()),
):
config_info = await get_vector_store_info(
data=VectorStoreInfoRequest(vector_store_id=self.CONFIG_ID), user_api_key_dict=self._admin()
)
db_info = await get_vector_store_info(
data=VectorStoreInfoRequest(vector_store_id=self.DB_ID), user_api_key_dict=self._admin()
)
assert config_info["vector_store"].is_config is True
assert db_info["vector_store"].is_config is False
@pytest.mark.asyncio
async def test_new_with_config_store_id_is_rejected_before_db_write(self):
prisma = MagicMock()
prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_managedvectorstorestable.create = AsyncMock()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam
patch.object(litellm, "vector_store_registry", self._registry()),
pytest.raises(HTTPException) as exc_info,
):
await new_vector_store(
vector_store={"vector_store_id": self.CONFIG_ID, "custom_llm_provider": "openai"},
user_api_key_dict=self._admin(),
)
assert exc_info.value.status_code == 400, exc_info.value.detail
assert exc_info.value.detail["vector_store_id"] == self.CONFIG_ID
assert "config file" in exc_info.value.detail["error"]
prisma.db.litellm_managedvectorstorestable.create.assert_not_called()
@pytest.mark.asyncio
async def test_update_of_config_store_is_rejected_before_db_write(self):
from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store
from litellm.types.vector_stores import VectorStoreUpdateRequest
prisma = MagicMock()
prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_managedvectorstorestable.update = AsyncMock()
registry = self._registry()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam
patch.object(litellm, "vector_store_registry", registry),
pytest.raises(HTTPException) as exc_info,
):
await update_vector_store(
data=VectorStoreUpdateRequest(vector_store_id=self.CONFIG_ID, vector_store_name="renamed"),
user_api_key_dict=self._admin(),
)
assert exc_info.value.status_code == 400, exc_info.value.detail
assert exc_info.value.detail["vector_store_id"] == self.CONFIG_ID
prisma.db.litellm_managedvectorstorestable.update.assert_not_called()
assert registry.get_litellm_managed_vector_store_from_registry(self.CONFIG_ID)["vector_store_name"] == "config-store"
@pytest.mark.asyncio
async def test_delete_of_config_store_is_rejected_and_store_stays_registered(self):
from litellm.proxy.vector_store_endpoints.management_endpoints import delete_vector_store
from litellm.types.vector_stores import VectorStoreDeleteRequest
prisma = MagicMock()
prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_managedvectorstorestable.delete = AsyncMock()
registry = self._registry()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam
patch.object(litellm, "vector_store_registry", registry),
pytest.raises(HTTPException) as exc_info,
):
await delete_vector_store(
data=VectorStoreDeleteRequest(vector_store_id=self.CONFIG_ID), user_api_key_dict=self._admin()
)
assert exc_info.value.status_code == 400, exc_info.value.detail
assert exc_info.value.detail["vector_store_id"] == self.CONFIG_ID
prisma.db.litellm_managedvectorstorestable.delete.assert_not_called()
assert registry.is_config_vector_store(self.CONFIG_ID) is True
@pytest.mark.asyncio
async def test_delete_of_db_store_still_works(self):
from litellm.proxy.vector_store_endpoints.management_endpoints import delete_vector_store
from litellm.types.vector_stores import VectorStoreDeleteRequest
row = MagicMock()
row.model_dump = MagicMock(return_value=self._db_row(self.DB_ID, "db-store"))
prisma = MagicMock()
prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=row)
prisma.db.litellm_managedvectorstorestable.delete = AsyncMock()
registry = self._registry()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam
patch.object(litellm, "vector_store_registry", registry),
):
response = await delete_vector_store(
data=VectorStoreDeleteRequest(vector_store_id=self.DB_ID), user_api_key_dict=self._admin()
)
assert response["status"] == "success", response
prisma.db.litellm_managedvectorstorestable.delete.assert_awaited_once_with(where={"vector_store_id": self.DB_ID})
assert registry.get_litellm_managed_vector_store_from_registry(self.DB_ID) is None

View file

@ -38,6 +38,7 @@ from litellm.types.utils import (
Usage,
)
from litellm.types.videos.main import VideoObject
from litellm.utils import supports_prompt_caching
@pytest.fixture
@ -4536,3 +4537,99 @@ def test_cost_per_token_bedrock_nemotron_super_3_uses_eu_west_2_entry_not_us_rat
assert prompt_usd == pytest.approx(prompt_tokens * regional["input_cost_per_token"])
assert completion_usd == pytest.approx(completion_tokens * regional["output_cost_per_token"])
GPT_REALTIME_2_FAMILY: Final = (
"azure/gpt-realtime-2.1",
"azure/gpt-realtime-2.1-mini",
"gpt-realtime-2",
"gpt-realtime-2.1",
"gpt-realtime-2.1-mini",
)
def test_gpt_realtime_2_family_prices_audio_cache_writes_and_reads_alike(_local_model_cost_map: None) -> None:
audio_cache_rates: Final = {
model: (
litellm.model_cost[model].get("cache_read_input_audio_token_cost"),
litellm.model_cost[model].get("cache_creation_input_audio_token_cost"),
)
for model in GPT_REALTIME_2_FAMILY
}
# Azure publishes one cached-audio meter per gpt-realtime-2 deployment,
# https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/, checked 2026-09-23
assert all(read is not None and write == read for read, write in audio_cache_rates.values()), audio_cache_rates
assert len(audio_cache_rates) == len(GPT_REALTIME_2_FAMILY)
GEMINI_LIVE_NATIVE_AUDIO_CASES: Final = (
("gemini-live-2.5-flash-native-audio", "vertex_ai"),
("gemini-live-2.5-flash-preview-native-audio-09-2025", "vertex_ai"),
("gemini/gemini-live-2.5-flash-preview-native-audio-09-2025", "gemini"),
)
@pytest.mark.parametrize(("model", "provider"), GEMINI_LIVE_NATIVE_AUDIO_CASES)
def test_gemini_live_native_audio_carries_no_cached_input_rate(
_local_model_cost_map: None, model: str, provider: str
) -> None:
# the Vertex pricing table prints N/A for cached input on every Live row,
# https://cloud.google.com/vertex-ai/generative-ai/pricing, checked 2026-09-23
assert litellm.get_model_info(model, custom_llm_provider=provider)["cache_read_input_token_cost"] is None
prompt_usd, _ = cost_per_token(
model=model,
prompt_tokens=101_000,
completion_tokens=0,
custom_llm_provider=provider,
usage_object=Usage(
prompt_tokens=101_000,
completion_tokens=0,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=100_000),
),
)
fresh_usd, _ = cost_per_token(
model=model,
prompt_tokens=101_000,
completion_tokens=0,
custom_llm_provider=provider,
usage_object=Usage(prompt_tokens=101_000, completion_tokens=0),
)
assert prompt_usd == pytest.approx(fresh_usd), (
"with no cached rate the cached tokens bill at the input rate, so a phantom discount cannot appear"
)
assert prompt_usd > 0
@pytest.mark.parametrize(("model", "provider"), GEMINI_LIVE_NATIVE_AUDIO_CASES)
def test_gemini_live_native_audio_declares_prompt_caching_unsupported(
_local_model_cost_map: None, model: str, provider: str
) -> None:
# the Vertex context-caching supported-model lists contain no Live model while 2.5 Flash is listed,
# https://cloud.google.com/vertex-ai/generative-ai/docs/context-cache/context-cache-overview, checked 2026-09-23
assert litellm.get_model_info(model, custom_llm_provider=provider)["supports_prompt_caching"] is False
assert supports_prompt_caching(model=model, custom_llm_provider=provider) is False
assert supports_prompt_caching(model="gemini-2.5-flash", custom_llm_provider="vertex_ai") is True, (
"control: the helper swallows a lookup error into False, so without this a broken lookup reads as a pass"
)
@pytest.mark.parametrize(
"model",
["gemini-live-2.5-flash-native-audio", "vertex_ai/gemini-live-2.5-flash-native-audio"],
)
def test_gemini_live_native_audio_limits_and_capabilities_match_vendor_model_card(
_local_model_cost_map: None, model: str
) -> None:
info = litellm.get_model_info(model)
# the Vertex model card for gemini-live-2.5-flash-native-audio publishes these limits and flags,
# https://cloud.google.com/vertex-ai/generative-ai/docs/models, checked 2026-09-23
assert info["max_input_tokens"] == 131072
assert info["max_output_tokens"] == 65536
assert info["max_tokens"] == 65536
assert info["supports_response_schema"] is False
assert info["supports_url_context"] is False
assert info["supports_pdf_input"] is False

View file

@ -8,7 +8,7 @@ from fastapi.testclient import TestClient
from datetime import datetime, timezone
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
import litellm
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
@ -182,3 +182,70 @@ def test_search_uses_registry_credentials():
assert getattr(called_params, "aws_region_name") == "us-east-1"
finally:
litellm.vector_store_registry = original_registry
def _config_registry(vector_store_id: str = "vs_from_config") -> VectorStoreRegistry:
registry = VectorStoreRegistry(vector_stores=[])
registry.load_vector_stores_from_config(
[
{
"vector_store_name": "config-store",
"litellm_params": {"vector_store_id": vector_store_id, "custom_llm_provider": "openai"},
}
]
)
return registry
def _db_store(vector_store_id: str, vector_store_name: str) -> LiteLLM_ManagedVectorStore:
return LiteLLM_ManagedVectorStore(
vector_store_id=vector_store_id,
custom_llm_provider="openai",
vector_store_name=vector_store_name,
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
def test_config_loaded_store_is_marked_config_owned_and_db_store_is_not():
registry = _config_registry()
registry.add_vector_store_to_registry(_db_store("vs_from_db", "db-store"))
assert registry.get_litellm_managed_vector_store_from_registry("vs_from_config")["is_config"] is True
assert registry.is_config_vector_store("vs_from_config") is True
assert registry.is_config_vector_store("vs_from_db") is False
assert registry.is_config_vector_store("vs_unknown") is False
def test_db_row_does_not_overwrite_config_owned_store_in_registry():
registry = _config_registry()
registry.add_vector_store_to_registry(_db_store("vs_from_db", "db-store"))
registry.update_vector_store_in_registry("vs_from_config", _db_store("vs_from_config", "renamed-in-db"))
registry.update_vector_store_in_registry("vs_from_db", _db_store("vs_from_db", "renamed-in-db"))
assert registry.get_litellm_managed_vector_store_from_registry("vs_from_config") == {
**registry.get_litellm_managed_vector_store_from_registry("vs_from_config"),
"vector_store_name": "config-store",
"is_config": True,
}
assert registry.get_litellm_managed_vector_store_from_registry("vs_from_db")["vector_store_name"] == "renamed-in-db"
@pytest.mark.asyncio
async def test_config_owned_store_survives_db_liveness_check_while_missing_db_store_is_evicted():
registry = _config_registry()
registry.add_vector_store_to_registry(_db_store("vs_from_db", "db-store"))
prisma_client = MagicMock()
prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None)
to_run = await registry.pop_vector_stores_to_run_with_db_fallback(
non_default_params={"vector_store_ids": ["vs_from_config", "vs_from_db"]},
prisma_client=prisma_client,
)
assert [vs["vector_store_id"] for vs in to_run] == ["vs_from_config"]
assert [vs["vector_store_id"] for vs in registry.vector_stores] == ["vs_from_config"]
prisma_client.db.litellm_managedvectorstorestable.find_unique.assert_awaited_once_with(
where={"vector_store_id": "vs_from_db"}
)

View file

@ -59,7 +59,16 @@ describe("VectorStoreTable", () => {
it("should render every column header", () => {
render(<VectorStoreTable {...defaultProps} />);
for (const header of ["Vector Store ID", "Name", "Description", "Files", "Provider", "Created At", "Updated At"]) {
for (const header of [
"Vector Store ID",
"Name",
"Description",
"Source",
"Files",
"Provider",
"Created At",
"Updated At",
]) {
expect(screen.getByText(header)).toBeInTheDocument();
}
});
@ -112,6 +121,36 @@ describe("VectorStoreTable", () => {
expect(mockOnDelete).toHaveBeenCalledWith("vs-newer");
});
it("should label each row's source as Config or DB", () => {
const configStore: VectorStore = { ...mockVectorStores[1], vector_store_id: "vs-config", is_config: true };
render(<VectorStoreTable {...defaultProps} data={[mockVectorStores[0], configStore]} />);
const rows = screen.getAllByRole("row").slice(1);
const dbRow = rows.find((row) => within(row).queryByText("vs-newer"));
const configRow = rows.find((row) => within(row).queryByText("vs-config"));
expect(within(dbRow!).getByText("DB")).toBeInTheDocument();
expect(within(dbRow!).queryByText("Config")).not.toBeInTheDocument();
expect(within(configRow!).getByText("Config")).toBeInTheDocument();
expect(within(configRow!).queryByText("DB")).not.toBeInTheDocument();
});
it("should keep edit and delete disabled for a config-defined store while copy still works", async () => {
const user = userEvent.setup();
const configStore: VectorStore = { ...mockVectorStores[1], vector_store_id: "vs-config", is_config: true };
render(<VectorStoreTable {...defaultProps} data={[mockVectorStores[0], configStore]} />);
await user.click(screen.getByTestId("vector-store-actions-vs-config"));
const editItem = await screen.findByTestId("vector-store-action-edit");
const deleteItem = screen.getByTestId("vector-store-action-delete");
expect(editItem).toHaveAttribute("aria-disabled", "true");
expect(deleteItem).toHaveAttribute("aria-disabled", "true");
expect(screen.getByText(/Read only: this vector store is defined in the config file/)).toBeVisible();
await user.click(editItem);
await user.click(deleteItem);
expect(mockOnEdit).not.toHaveBeenCalled();
expect(mockOnDelete).not.toHaveBeenCalled();
await user.click(screen.getByTestId("vector-store-action-copy"));
expect(await window.navigator.clipboard.readText()).toBe("vs-config");
});
it("should copy the vector store ID through the actions menu", async () => {
const user = userEvent.setup();
render(<VectorStoreTable {...defaultProps} />);
@ -119,4 +158,12 @@ describe("VectorStoreTable", () => {
await user.click(await screen.findByTestId("vector-store-action-copy"));
expect(await window.navigator.clipboard.readText()).toBe("vs-newer");
});
it("should not show the read-only hint for a database-backed store", async () => {
const user = userEvent.setup();
render(<VectorStoreTable {...defaultProps} />);
await user.click(screen.getByTestId("vector-store-actions-vs-newer"));
await screen.findByTestId("vector-store-action-edit");
expect(screen.queryByText(/Read only: this vector store is defined in the config file/)).not.toBeInTheDocument();
});
});

View file

@ -4,7 +4,7 @@ import { ColumnDef } from "@tanstack/react-table";
import { Copy, MoreHorizontal, Pencil, Trash2 } from "lucide-react";
import { DataTableSortHeader } from "@/components/shared/DataTable";
import { CellTooltip, DateCell, IdentityCell } from "@/components/shared/table_cells";
import { CellTooltip, DateCell, IdentityCell, StatusBadge } from "@/components/shared/table_cells";
import { getVectorStoreProviderLogoAndName } from "@/components/vector_store_providers";
import { buttonVariants } from "@/components/ui/button";
import {
@ -18,6 +18,9 @@ import { VectorStore } from "@/components/vector_store_management/types";
import { cn } from "@/lib/cva.config";
import { copyToClipboard } from "@/utils/dataUtils";
const CONFIG_STORE_HINT =
"Read only: this vector store is defined in the config file and cannot be edited or deleted on the dashboard.";
function VectorStoreProviderCell({ provider }: { provider: string }) {
const { displayName, logo } = getVectorStoreProviderLogoAndName(provider);
return (
@ -64,6 +67,7 @@ interface VectorStoreRowActionsProps {
}
function VectorStoreRowActions({ vectorStore, onEdit, onDelete }: VectorStoreRowActionsProps) {
const isFromConfig = vectorStore.is_config ?? false;
return (
<DropdownMenu>
<DropdownMenuTrigger
@ -74,7 +78,11 @@ function VectorStoreRowActions({ vectorStore, onEdit, onDelete }: VectorStoreRow
<MoreHorizontal className="size-4" />
</DropdownMenuTrigger>
<DropdownMenuContent align="end" className="w-52">
<DropdownMenuItem data-testid="vector-store-action-edit" onClick={() => onEdit(vectorStore.vector_store_id)}>
<DropdownMenuItem
data-testid="vector-store-action-edit"
disabled={isFromConfig}
onClick={() => onEdit(vectorStore.vector_store_id)}
>
<Pencil />
Edit
</DropdownMenuItem>
@ -89,11 +97,17 @@ function VectorStoreRowActions({ vectorStore, onEdit, onDelete }: VectorStoreRow
<DropdownMenuItem
variant="destructive"
data-testid="vector-store-action-delete"
disabled={isFromConfig}
onClick={() => onDelete(vectorStore.vector_store_id)}
>
<Trash2 />
Delete
</DropdownMenuItem>
{isFromConfig && (
<div data-testid="vector-store-config-hint" className="px-2 py-1.5 text-xs text-muted-foreground">
{CONFIG_STORE_HINT}
</div>
)}
</DropdownMenuContent>
</DropdownMenu>
);
@ -158,6 +172,18 @@ export const getVectorStoreTableColumns = ({
);
},
},
{
id: "source",
accessorFn: (row) => row.is_config ?? false,
meta: { title: "Source", skeleton: "badge" },
header: ({ column }) => <DataTableSortHeader column={column} title="Source" />,
size: 110,
enableSorting: true,
cell: ({ row }) => {
const isFromConfig = row.original.is_config ?? false;
return <StatusBadge tone={isFromConfig ? "neutral" : "info"} label={isFromConfig ? "Config" : "DB"} />;
},
},
{
id: "files",
meta: { title: "Files" },

View file

@ -59,6 +59,61 @@ describe("VectorStoreInfoView", () => {
expect(await screen.findByText("Vector Store ID: vs-1")).toBeInTheDocument();
});
it("should render a config-defined store read-only for an admin, even when opened in edit mode", async () => {
mockVectorStoreInfoCall.mockResolvedValue({
vector_store: {
vector_store_id: "vs-config",
vector_store_name: "config-store",
custom_llm_provider: "openai",
created_at: "2024-01-01T00:00:00Z",
updated_at: "2024-01-01T00:00:00Z",
is_config: true,
},
});
render(
<VectorStoreInfoView
vectorStoreId="vs-config"
onClose={vi.fn()}
accessToken="sk-test"
is_admin={true}
editVectorStore={true}
/>,
);
expect(await screen.findByText("Vector Store ID: vs-config")).toBeInTheDocument();
expect(screen.getByText("Read only: defined in the config file")).toBeInTheDocument();
expect(screen.getByText("Config")).toBeInTheDocument();
expect(screen.queryByText("DB")).not.toBeInTheDocument();
expect(screen.getByText("Vector Store Details")).toBeInTheDocument();
expect(screen.queryByRole("button", { name: "Edit Vector Store" })).not.toBeInTheDocument();
expect(screen.queryByRole("button", { name: /Save/ })).not.toBeInTheDocument();
});
it("should still offer editing for a database-backed store", async () => {
mockVectorStoreInfoCall.mockResolvedValue({
vector_store: {
vector_store_id: "vs-db",
vector_store_name: "db-store",
custom_llm_provider: "openai",
created_at: "2024-01-01T00:00:00Z",
updated_at: "2024-01-01T00:00:00Z",
is_config: false,
},
});
render(
<VectorStoreInfoView
vectorStoreId="vs-db"
onClose={vi.fn()}
accessToken="sk-test"
is_admin={true}
editVectorStore={false}
/>,
);
expect(await screen.findByText("Vector Store ID: vs-db")).toBeInTheDocument();
expect(screen.queryByText("Read only: defined in the config file")).not.toBeInTheDocument();
expect(screen.getByText("DB")).toBeInTheDocument();
expect(screen.getAllByRole("button", { name: "Edit Vector Store" }).length).toBeGreaterThan(0);
});
it("should show a not-found state with a working back button when the fetch fails instead of loading forever", async () => {
const user = userEvent.setup();
const onClose = vi.fn();

View file

@ -1,5 +1,5 @@
import React, { useState, useEffect } from "react";
import { ArrowLeft, CircleHelp } from "lucide-react";
import { ArrowLeft, CircleHelp, Lock } from "lucide-react";
import { z } from "zod/v4";
import {
vectorStoreInfoCall,
@ -15,6 +15,8 @@ import VectorStoreTester from "./VectorStoreTester";
import { toast } from "@/lib/toast";
import { FieldGroup } from "@/components/ui/field";
import { FormField } from "@/components/shared/form/FormField";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { StatusBadge } from "@/components/shared/table_cells";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Card, CardContent } from "@/components/ui/card";
@ -200,6 +202,9 @@ const VectorStoreInfoView: React.FC<VectorStoreInfoViewProps> = ({
return <div>Loading...</div>;
}
const canEdit = is_admin && !vectorStoreDetails.is_config;
const showEditForm = isEditing && canEdit;
return (
<div className="p-4 max-w-full">
<div className="flex justify-between items-center mb-6">
@ -208,14 +213,31 @@ const VectorStoreInfoView: React.FC<VectorStoreInfoViewProps> = ({
<ArrowLeft />
Back to Vector Stores
</Button>
<h1 className="text-xl font-semibold">Vector Store ID: {vectorStoreDetails.vector_store_id}</h1>
<div className="flex items-center gap-2">
<h1 className="text-xl font-semibold">Vector Store ID: {vectorStoreDetails.vector_store_id}</h1>
<StatusBadge
tone={vectorStoreDetails.is_config ? "neutral" : "info"}
label={vectorStoreDetails.is_config ? "Config" : "DB"}
/>
</div>
<p className="text-sm text-muted-foreground">
{vectorStoreDetails.vector_store_description || "No description"}
</p>
</div>
{is_admin && !isEditing && <Button onClick={startEditing}>Edit Vector Store</Button>}
{canEdit && !isEditing && <Button onClick={startEditing}>Edit Vector Store</Button>}
</div>
{vectorStoreDetails.is_config && (
<Alert variant="info" className="mb-4">
<Lock className="size-4" aria-hidden />
<AlertTitle>Read only: defined in the config file</AlertTitle>
<AlertDescription>
This vector store comes from the proxy config YAML, so it cannot be edited or deleted on the dashboard.
Change or remove it in the config file and restart the proxy.
</AlertDescription>
</Alert>
)}
<Tabs defaultValue="details">
<TabsList variant="line" className="mb-6 h-auto w-full justify-start rounded-none p-0">
<TabsTrigger value="details" className="flex-none rounded-none px-4 py-2">
@ -227,7 +249,7 @@ const VectorStoreInfoView: React.FC<VectorStoreInfoViewProps> = ({
</TabsList>
<TabsContent value="details" keepMounted>
{isEditing ? (
{showEditForm ? (
<div>
<div className="flex justify-between items-center mb-4">
<h3 className="text-lg font-medium">Edit Vector Store</h3>
@ -373,7 +395,7 @@ const VectorStoreInfoView: React.FC<VectorStoreInfoViewProps> = ({
<div>
<div className="flex justify-between items-center mb-4">
<h3 className="text-lg font-medium">Vector Store Details</h3>
{is_admin && <Button onClick={startEditing}>Edit Vector Store</Button>}
{canEdit && <Button onClick={startEditing}>Edit Vector Store</Button>}
</div>
<Card>
<CardContent>

Some files were not shown because too many files have changed in this diff Show more