mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
Merge remote-tracking branch 'origin/main' into litellm_anthropic_guardrail_system_and_tool_use
main's test_one_text_per_row_over_a_system_prompt_is_rejected_by_name assumed the top-level system prompt stays out of the scanned texts. This branch scans it, so one text per structured row now lines up and the rewrite is applied; the test asserts that, and a multi-block system prompt keeps the length-guard rejection covered.
This commit is contained in:
commit
04410967ff
66 changed files with 3679 additions and 372 deletions
|
|
@ -2915,6 +2915,25 @@ jobs:
|
|||
exit 1
|
||||
fi
|
||||
|
||||
provider_replay_harness:
|
||||
docker:
|
||||
- *python312_image
|
||||
working_directory: ~/project
|
||||
resource_class: medium
|
||||
steps:
|
||||
- setup_litellm_test_deps
|
||||
- run:
|
||||
name: Test provider replay harness
|
||||
command: |
|
||||
mkdir -p test-results/provider-replay-harness
|
||||
uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures \
|
||||
--junitxml=test-results/provider-replay-harness/junit.xml \
|
||||
tests/e2e/test_provider_edge.py tests/e2e/test_fixture_bundle.py \
|
||||
tests/e2e/test_fixture_canonical.py tests/e2e/test_fixture_mode.py \
|
||||
tests/code_coverage_tests/test_provider_replay_harness.py
|
||||
- store_test_results:
|
||||
path: test-results/provider-replay-harness
|
||||
|
||||
integration_contracts:
|
||||
parameters:
|
||||
suite:
|
||||
|
|
@ -2967,6 +2986,7 @@ workflows:
|
|||
only:
|
||||
- main
|
||||
- /litellm_.*/
|
||||
- provider_replay_harness
|
||||
- base_sdk_install:
|
||||
filters: *main_branches
|
||||
- local_testing_part1:
|
||||
|
|
|
|||
15
.github/e2e-stack/assert_tests_ran.py
vendored
15
.github/e2e-stack/assert_tests_ran.py
vendored
|
|
@ -3,6 +3,9 @@ import xml.etree.ElementTree as ET
|
|||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "tests/e2e"))
|
||||
from coverage_registry.management_cases import MANAGEMENT_CASES
|
||||
|
||||
|
||||
def main() -> int:
|
||||
selected: Final = tuple(sys.argv[2:])
|
||||
|
|
@ -16,6 +19,17 @@ def main() -> int:
|
|||
case.get("file") for case in cases if all(case.find(tag) is None for tag in ("skipped", "failure", "error"))
|
||||
)
|
||||
missing: Final = tuple(path for path in selected if path not in passed)
|
||||
required_nodes: Final = frozenset(case.node for case in MANAGEMENT_CASES if case.node.split("::", 1)[0] in selected)
|
||||
passed_nodes: Final = frozenset(
|
||||
prop.get("value")
|
||||
for case in cases
|
||||
if all(case.find(tag) is None for tag in ("skipped", "failure", "error"))
|
||||
for prop in case.findall("./properties/property")
|
||||
if prop.get("name") == "management_node"
|
||||
)
|
||||
missing_nodes: Final = required_nodes - passed_nodes
|
||||
for node in sorted(missing_nodes):
|
||||
_ = sys.stdout.write(f"::error::required management case did not pass: {node}\n")
|
||||
for path in selected:
|
||||
collected: Final = sum(case.get("file") == path for case in cases)
|
||||
skipped: Final = sum(case.get("file") == path and case.find("skipped") is not None for case in cases)
|
||||
|
|
@ -27,6 +41,7 @@ def main() -> int:
|
|||
if (
|
||||
selected
|
||||
and not missing
|
||||
and not missing_nodes
|
||||
and not any(case.find(tag) is not None for case in cases for tag in ("failure", "error"))
|
||||
):
|
||||
return 0
|
||||
|
|
|
|||
5
.github/e2e-stack/oidc-profile.sh
vendored
Executable file
5
.github/e2e-stack/oidc-profile.sh
vendored
Executable file
|
|
@ -0,0 +1,5 @@
|
|||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
|
||||
cd "${REPO_ROOT}"
|
||||
exec uv run --no-sync python tests/e2e/idp.py "$@"
|
||||
2
.github/e2e-stack/select_tests.py
vendored
2
.github/e2e-stack/select_tests.py
vendored
|
|
@ -12,6 +12,8 @@ UNSUPPORTED: Final = re.compile(
|
|||
HARNESS: Final = re.compile(
|
||||
r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$"
|
||||
r"|^tests/e2e/idp_realm\.json$"
|
||||
r"|^tests/e2e/management/(management_client|jwt_actors|conftest)\.py$"
|
||||
r"|^tests/e2e/coverage_registry/management_cases\.py$"
|
||||
r"|^tests/e2e/gateway/"
|
||||
r"|^\.github/e2e-stack/"
|
||||
r"|^\.github/workflows/test-e2e-changed\.yml$"
|
||||
|
|
|
|||
2
.github/workflows/_test-unit-base.yml
vendored
2
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -57,6 +57,7 @@ permissions:
|
|||
|
||||
env:
|
||||
UV_PYTHON: "3.12"
|
||||
LITELLM_LOCAL_MODEL_COST_MAP: "True"
|
||||
|
||||
jobs:
|
||||
run:
|
||||
|
|
@ -113,6 +114,7 @@ jobs:
|
|||
if: steps.changes.outputs.decision != 'skip'
|
||||
timeout-minutes: 8
|
||||
run: |
|
||||
diff -u model_prices_and_context_window.json litellm/model_prices_and_context_window_backup.json
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
|
||||
uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]'
|
||||
|
||||
|
|
|
|||
2
.github/workflows/test-e2e-changed.yml
vendored
2
.github/workflows/test-e2e-changed.yml
vendored
|
|
@ -183,7 +183,7 @@ jobs:
|
|||
log="${RUNNER_TEMP}/e2e-pass-${pass}.log"
|
||||
echo "::group::pass ${pass} of 3"
|
||||
set +e
|
||||
uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v -p no:cacheprovider \
|
||||
uv run --no-sync pytest "${test_files[@]}" --rootdir=. -v --reruns 0 -p no:cacheprovider \
|
||||
-o junit_family=xunit1 --junitxml="${report}" > "${log}" 2>&1
|
||||
status=$?
|
||||
uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}"
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
|
||||
from litellm.llms.vertex_ai.batches.transformation import vertex_prompt_tokens_details
|
||||
from litellm.types.llms.openai import Batch
|
||||
from litellm.types.utils import ModelInfo, Usage
|
||||
from litellm.utils import token_counter
|
||||
|
|
@ -356,6 +357,7 @@ def calculate_vertex_ai_batch_cost_and_usage(
|
|||
prompt_tokens=_prompt,
|
||||
completion_tokens=_completion,
|
||||
total_tokens=_total,
|
||||
prompt_tokens_details=vertex_prompt_tokens_details(usage_metadata),
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -205,21 +205,35 @@ def _extract_anthropic_tool_exchange_spans(
|
|||
return spans, None
|
||||
|
||||
|
||||
def _message_has_cache_control(message: Mapping[str, object]) -> bool:
|
||||
if message.get("cache_control") is not None:
|
||||
return True
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, list):
|
||||
return any(isinstance(part, Mapping) and part.get("cache_control") is not None for part in content)
|
||||
return False
|
||||
|
||||
|
||||
def get_protected_indices(messages: Sequence[Mapping[str, object]]) -> tuple[int, ...]:
|
||||
"""
|
||||
Return indices of messages that must never be compressed:
|
||||
- All system messages
|
||||
- The last user message
|
||||
- The last assistant message
|
||||
- Any message carrying an Anthropic cache_control breakpoint
|
||||
|
||||
The last user message is what the model is being asked to act on right now,
|
||||
so compressing it replaces the live instruction with a marker. Compression
|
||||
guardrails share this policy; see the Headroom guardrail.
|
||||
guardrails share this policy; see the Headroom guardrail. A cache_control
|
||||
breakpoint pins the provider's prompt-cache prefix to that row's exact
|
||||
bytes, so rewriting a marked row anywhere in history turns the next
|
||||
request's cache read into a cache write.
|
||||
"""
|
||||
system_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system")
|
||||
last_user: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:]
|
||||
last_assistant = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")[-1:]
|
||||
return system_indices + last_user + last_assistant
|
||||
assistant_indices: Final = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")
|
||||
cache_control_indices: Final = tuple(index for index, msg in enumerate(messages) if _message_has_cache_control(msg))
|
||||
return tuple(dict.fromkeys(system_indices + last_user + assistant_indices[-1:] + cache_control_indices))
|
||||
|
||||
|
||||
def _combine_scores(
|
||||
|
|
@ -421,7 +435,7 @@ def compress(
|
|||
combined_scores = bm25_scores
|
||||
|
||||
# Protected messages are never compressed
|
||||
protected_indices: Final = get_protected_indices(normalized_messages)
|
||||
protected_indices: Final = get_protected_indices(original_messages)
|
||||
kept_indices: set[int] = set(protected_indices)
|
||||
|
||||
tool_exchange_spans: list[set[int]] = []
|
||||
|
|
|
|||
|
|
@ -2278,6 +2278,19 @@ def default_video_cost_calculator(
|
|||
return 0.0
|
||||
|
||||
|
||||
def _batch_rate(
|
||||
model_info: ModelInfo,
|
||||
key: Literal[
|
||||
"input_cost_per_audio_token_batches",
|
||||
"input_cost_per_image_token_batches",
|
||||
"input_cost_per_video_token_batches",
|
||||
],
|
||||
fallback: float,
|
||||
) -> float:
|
||||
rate: Final = model_info.get(key)
|
||||
return fallback if rate is None else rate
|
||||
|
||||
|
||||
def batch_cost_calculator(
|
||||
usage: Usage,
|
||||
model: str,
|
||||
|
|
@ -2337,7 +2350,29 @@ def batch_cost_calculator(
|
|||
total_prompt_cost = 0.0
|
||||
total_completion_cost = 0.0
|
||||
if input_cost_per_token_batches is not None:
|
||||
total_prompt_cost = usage.prompt_tokens * input_cost_per_token_batches
|
||||
batch_details: Final = parse_prompt_tokens_details(usage)
|
||||
audio_tokens, image_tokens, video_tokens = (
|
||||
batch_details["audio_tokens"],
|
||||
batch_details["image_tokens"],
|
||||
batch_details["video_tokens"],
|
||||
)
|
||||
modality_rates: Final = (
|
||||
_batch_rate(model_info, "input_cost_per_audio_token_batches", input_cost_per_token_batches),
|
||||
_batch_rate(model_info, "input_cost_per_image_token_batches", input_cost_per_token_batches),
|
||||
_batch_rate(model_info, "input_cost_per_video_token_batches", input_cost_per_token_batches),
|
||||
)
|
||||
total_prompt_cost = sum(
|
||||
tokens * rate
|
||||
for tokens, rate in zip(
|
||||
(
|
||||
max((usage.prompt_tokens or 0) - audio_tokens - image_tokens - video_tokens, 0),
|
||||
audio_tokens,
|
||||
image_tokens,
|
||||
video_tokens,
|
||||
),
|
||||
(input_cost_per_token_batches, *modality_rates),
|
||||
)
|
||||
)
|
||||
elif input_cost_per_token:
|
||||
details: Final = parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens: Final = details["cache_hit_tokens"]
|
||||
|
|
|
|||
|
|
@ -956,12 +956,16 @@ def _calculate_input_cost(
|
|||
)
|
||||
|
||||
### AUDIO COST
|
||||
if prompt_tokens_details["audio_tokens"]:
|
||||
if prompt_tokens_details["audio_tokens"] and not (
|
||||
prompt_tokens_details["audio_length_seconds"] and model_info.get("input_cost_per_audio_per_second") is not None
|
||||
):
|
||||
audio_cost_key: Final = _get_service_tier_cost_key("input_cost_per_audio_token", service_tier)
|
||||
prompt_cost += calculate_cost_component(model_info, audio_cost_key, prompt_tokens_details["audio_tokens"])
|
||||
|
||||
### IMAGE TOKEN COST
|
||||
if prompt_tokens_details["image_tokens"]:
|
||||
if prompt_tokens_details["image_tokens"] and not (
|
||||
prompt_tokens_details["image_count"] and model_info.get("input_cost_per_image") is not None
|
||||
):
|
||||
# For image token costs:
|
||||
# First check if input_cost_per_image_token is available. If not, default to generic input_cost_per_token.
|
||||
image_token_cost_key = "input_cost_per_image_token"
|
||||
|
|
@ -970,7 +974,9 @@ def _calculate_input_cost(
|
|||
prompt_cost += calculate_cost_component(model_info, image_token_cost_key, prompt_tokens_details["image_tokens"])
|
||||
|
||||
### VIDEO TOKEN COST
|
||||
if prompt_tokens_details["video_tokens"]:
|
||||
if prompt_tokens_details["video_tokens"] and not (
|
||||
prompt_tokens_details["video_length_seconds"] and model_info.get("input_cost_per_video_per_second") is not None
|
||||
):
|
||||
video_token_cost_key = "input_cost_per_video_token"
|
||||
if model_info.get(video_token_cost_key) is None:
|
||||
video_token_cost_key = "input_cost_per_token"
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
scoped_structured_message_indices,
|
||||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
unappliable_request_rewrite,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
|
|
@ -674,6 +675,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
preserve_system_messages=has_midturn_system_message,
|
||||
)
|
||||
else:
|
||||
if guardrailed_texts and len(guardrailed_texts) != len(scanned):
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
data=data,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable, Iterator, Sequence
|
||||
from typing import Final, TypeVar
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from typing import Final, TypeVar, cast # noqa: TID251 # a rebuilt chat row has no typed constructor across roles
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -364,3 +364,67 @@ def merge_guardrailed_scoped_messages(
|
|||
yield from appended
|
||||
|
||||
return list(_merged())
|
||||
|
||||
|
||||
def _content_part_text(part: object) -> str | None:
|
||||
if not isinstance(part, Mapping):
|
||||
return None
|
||||
text: Final = part.get("text")
|
||||
return text if isinstance(text, str) else None
|
||||
|
||||
|
||||
def message_slot_texts(message: Mapping[str, object]) -> tuple[str, ...]:
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, str):
|
||||
return (content,)
|
||||
if isinstance(content, list):
|
||||
return tuple(text for part in content if (text := _content_part_text(part)) is not None)
|
||||
return ()
|
||||
|
||||
|
||||
def message_text_slot_count(message: AllMessageValues) -> int:
|
||||
return len(message_slot_texts(message))
|
||||
|
||||
|
||||
def _part_with_text(part: object, text: str) -> object:
|
||||
if not isinstance(part, Mapping):
|
||||
return part
|
||||
return {**part, "text": text} # mutable-ok: content parts stay JSON-plain dicts
|
||||
|
||||
|
||||
def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> Sequence[object]:
|
||||
remaining_texts: Final = iter(texts)
|
||||
return [ # mutable-ok: message content stays a JSON list
|
||||
_part_with_text(part, next(remaining_texts)) if _content_part_text(part) is not None else part
|
||||
for part in content
|
||||
]
|
||||
|
||||
|
||||
def message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues | None:
|
||||
"""Swap one rewritten text into each text slot of a chat row, in order.
|
||||
|
||||
A slot is a string ``content`` or one list part carrying a string ``text``;
|
||||
images and other parts ride along untouched. Returns None unless the counts
|
||||
line up exactly, so a rewrite never lands on the wrong slot.
|
||||
"""
|
||||
if message_text_slot_count(message) != len(texts):
|
||||
return None
|
||||
content: Final = message.get("content")
|
||||
if not isinstance(content, (str, list)):
|
||||
return message
|
||||
rewritten_content: Final = texts[0] if isinstance(content, str) else _content_with_slot_texts(content, texts)
|
||||
rewritten: Final = {**message, "content": rewritten_content} # mutable-ok: chat rows stay JSON-plain dicts
|
||||
return cast("AllMessageValues", rewritten) # cast-ok: the same row with only its text slots swapped
|
||||
|
||||
|
||||
class UnappliableRequestRewrite(Exception):
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
super().__init__(
|
||||
f"Guardrail '{guardrail_name}' rewrote the request in a way this endpoint cannot apply, "
|
||||
"so the request was rejected rather than sent unrewritten"
|
||||
)
|
||||
self.guardrail_name: Final = guardrail_name
|
||||
|
||||
|
||||
def unappliable_request_rewrite(guardrail_name: str | None) -> UnappliableRequestRewrite:
|
||||
return UnappliableRequestRewrite(guardrail_name or "unknown")
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
stream_item_items,
|
||||
unappliable_request_rewrite,
|
||||
)
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
|
|
@ -196,6 +197,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
else:
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
if guardrailed_texts and texts_to_check:
|
||||
if len(guardrailed_texts) != len(text_task_mappings):
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
await self._apply_guardrail_responses_to_input_texts(
|
||||
messages=messages,
|
||||
responses=guardrailed_texts,
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
stream_item_items,
|
||||
unappliable_request_rewrite,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
|
|
@ -495,13 +496,13 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
|
||||
elif isinstance(input_data, str):
|
||||
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(guardrailed_texts) > 1:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
|
||||
else:
|
||||
rewritten_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(rewritten_texts) != len(extracted.task_mappings):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite
|
||||
|
||||
raise UnappliableRequestRewrite(guardrail_to_apply.guardrail_name or "unknown")
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=rewritten_texts,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
from urllib.parse import unquote
|
||||
|
||||
|
|
@ -8,7 +9,36 @@ from litellm.llms.vertex_ai.common_utils import (
|
|||
)
|
||||
from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest
|
||||
from litellm.types.llms.vertex_ai import *
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
from litellm.types.utils import LiteLLMBatch, PromptTokensDetailsWrapper
|
||||
|
||||
|
||||
def vertex_prompt_tokens_details(
|
||||
usage_metadata: Mapping[str, object],
|
||||
) -> PromptTokensDetailsWrapper | None:
|
||||
raw_details: Final = usage_metadata.get("promptTokensDetails")
|
||||
if not isinstance(raw_details, list):
|
||||
return None
|
||||
|
||||
def _normalize(detail: object) -> tuple[str, int] | None:
|
||||
if not isinstance(detail, Mapping):
|
||||
return None
|
||||
modality: Final = detail.get("modality")
|
||||
token_count: Final = detail.get("tokenCount")
|
||||
if not isinstance(modality, str) or not isinstance(token_count, int):
|
||||
return None
|
||||
return modality.upper(), token_count
|
||||
|
||||
parsed_details: Final = tuple(_normalize(detail) for detail in raw_details)
|
||||
normalized: Final = tuple(detail for detail in parsed_details if detail is not None)
|
||||
if len(normalized) != len(parsed_details):
|
||||
return None
|
||||
|
||||
return PromptTokensDetailsWrapper(
|
||||
text_tokens=sum(token_count for modality, token_count in normalized if modality in ("TEXT", "DOCUMENT")),
|
||||
audio_tokens=sum(token_count for modality, token_count in normalized if modality == "AUDIO"),
|
||||
image_tokens=sum(token_count for modality, token_count in normalized if modality == "IMAGE"),
|
||||
video_tokens=sum(token_count for modality, token_count in normalized if modality == "VIDEO"),
|
||||
)
|
||||
|
||||
|
||||
class VertexAIBatchTransformation:
|
||||
|
|
|
|||
|
|
@ -298,8 +298,6 @@ def transform_openai_input_gemini_embed_content(
|
|||
|
||||
|
||||
_IMAGE_MIME_TYPES: Final = frozenset({"image/png", "image/jpeg"})
|
||||
_VIDEO_TOKENS_PER_SECOND: Final = 258.0
|
||||
_AUDIO_TOKENS_PER_SECOND: Final = 32.0
|
||||
_usage_metadata_adapter: Final = TypeAdapter(UsageMetadata)
|
||||
|
||||
|
||||
|
|
@ -339,11 +337,12 @@ def _is_image_element(
|
|||
return False
|
||||
|
||||
|
||||
def _count_input_images(
|
||||
def _is_image_only_input(
|
||||
input: GeminiEmbeddingInput,
|
||||
resolved_files: Mapping[str, Mapping[str, str]],
|
||||
) -> int:
|
||||
return sum(1 for element in _flatten_input(input) if _is_image_element(element, resolved_files))
|
||||
) -> bool:
|
||||
elements: Final = _flatten_input(input)
|
||||
return bool(elements) and all(_is_image_element(element, resolved_files) for element in elements)
|
||||
|
||||
|
||||
def _tokens_for_modality(details: Sequence[PromptTokensDetails], modality: str) -> int:
|
||||
|
|
@ -372,30 +371,29 @@ def _usage_from_embed_content_response(
|
|||
total_tokens: Final = usage_metadata.get("totalTokenCount") or prompt_tokens
|
||||
|
||||
details: Final[Sequence[PromptTokensDetails]] = usage_metadata.get("promptTokensDetails") or ()
|
||||
if not details:
|
||||
return Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
total_tokens=total_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=0,
|
||||
image_tokens=prompt_tokens if _is_image_only_input(input, resolved_files) else 0,
|
||||
),
|
||||
)
|
||||
|
||||
text_tokens: Final = _tokens_for_modality(details, "TEXT")
|
||||
audio_tokens: Final = _tokens_for_modality(details, "AUDIO")
|
||||
image_tokens: Final = _tokens_for_modality(details, "IMAGE")
|
||||
video_tokens: Final = _tokens_for_modality(details, "VIDEO")
|
||||
image_count: Final = _count_input_images(input, resolved_files)
|
||||
|
||||
video_length_seconds: Final = video_tokens / _VIDEO_TOKENS_PER_SECOND if video_tokens > 0 else 0.0
|
||||
audio_length_seconds: Final = audio_tokens / _AUDIO_TOKENS_PER_SECOND if audio_tokens > 0 else 0.0
|
||||
|
||||
# generic_cost_per_token rewrites text_tokens to the full prompt minus
|
||||
# other modalities when both text_tokens and image_count are zero. For
|
||||
# video, that misallocates video tokens to text; a 1-token floor sidesteps
|
||||
# the rewrite and keeps billing on input_cost_per_video_per_second.
|
||||
needs_video_text_floor: Final = video_length_seconds > 0 and text_tokens == 0 and image_count == 0
|
||||
resolved_text_tokens: Final = 1 if needs_video_text_floor else text_tokens
|
||||
|
||||
return Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
total_tokens=total_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=resolved_text_tokens,
|
||||
text_tokens=text_tokens,
|
||||
audio_tokens=audio_tokens,
|
||||
image_count=image_count,
|
||||
video_length_seconds=video_length_seconds,
|
||||
audio_length_seconds=audio_length_seconds,
|
||||
image_tokens=image_tokens,
|
||||
video_tokens=video_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -415,8 +413,7 @@ def process_embed_content_response(
|
|||
model_response: EmbeddingResponse to populate
|
||||
model: Model name
|
||||
response_json: Raw JSON response from embedContent endpoint
|
||||
resolved_files: Mapping of file references (files/abc) to {mime_type, uri},
|
||||
used to bill resolved image references at the per-image rate
|
||||
resolved_files: Mapping of file references to resolved metadata
|
||||
|
||||
Returns:
|
||||
EmbeddingResponse with single embedding
|
||||
|
|
|
|||
|
|
@ -25601,10 +25601,14 @@
|
|||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models"
|
||||
},
|
||||
"gemini-embedding-2-preview": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25615,13 +25619,14 @@
|
|||
"uses_embed_content": true
|
||||
},
|
||||
"gemini-embedding-2": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25633,10 +25638,14 @@
|
|||
"uses_embed_content": true
|
||||
},
|
||||
"vertex_ai/gemini-embedding-2-preview": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25648,13 +25657,14 @@
|
|||
"uses_embed_content": true
|
||||
},
|
||||
"vertex_ai/gemini-embedding-2": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25693,10 +25703,14 @@
|
|||
},
|
||||
"gemini/gemini-embedding-2-preview": {
|
||||
"deprecation_date": "2026-08-10",
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25709,10 +25723,14 @@
|
|||
"tpm": 10000000
|
||||
},
|
||||
"gemini/gemini-embedding-2": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@
|
|||
|
||||
import fnmatch
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
|
|
@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import ChatCompletionToolParam
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIMetadata,
|
||||
GenericGuardrailAPIRequest,
|
||||
|
|
@ -150,6 +150,26 @@ def _extract_inbound_headers(
|
|||
return None
|
||||
|
||||
|
||||
def _structured_rows_to_write_back(
|
||||
original_rows: Sequence[AllMessageValues] | None,
|
||||
shown_rows: Sequence[AllMessageValues] | None,
|
||||
returned_rows: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...] | None:
|
||||
"""The request model drops row keys its message types do not declare, so a
|
||||
row the server echoes back verbatim is restored to the original row object.
|
||||
A server that echoes every row back unchanged has not rewritten anything
|
||||
per row, so its answer is read from texts, as it was before rows could be
|
||||
returned at all."""
|
||||
if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows):
|
||||
return tuple(returned_rows)
|
||||
if all(returned == shown for shown, returned in zip(shown_rows, returned_rows)):
|
||||
return None
|
||||
return tuple(
|
||||
original if returned == shown else returned
|
||||
for original, shown, returned in zip(original_rows, shown_rows, returned_rows)
|
||||
)
|
||||
|
||||
|
||||
class GenericGuardrailAPI(CustomGuardrail):
|
||||
"""
|
||||
Generic Guardrail API integration for LiteLLM.
|
||||
|
|
@ -322,6 +342,8 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
texts: list,
|
||||
images: list[str] | None,
|
||||
tools: list[ChatCompletionToolParam] | None,
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
shown_messages: Sequence[AllMessageValues] | None,
|
||||
guardrail_response: GenericGuardrailAPIResponse,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
# Action is NONE or no modifications needed
|
||||
|
|
@ -336,6 +358,13 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
return_inputs["tools"] = guardrail_response.tools
|
||||
elif tools:
|
||||
return_inputs["tools"] = tools
|
||||
rows_to_write_back: Final = (
|
||||
_structured_rows_to_write_back(structured_messages, shown_messages, guardrail_response.structured_messages)
|
||||
if guardrail_response.structured_messages
|
||||
else None
|
||||
)
|
||||
if rows_to_write_back is not None:
|
||||
return_inputs["structured_messages"] = list(rows_to_write_back) # mutable-ok: guardrail inputs take a list
|
||||
if guardrail_response.stream_holdback_chars is not None:
|
||||
return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars
|
||||
return return_inputs
|
||||
|
|
@ -473,6 +502,8 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
texts=texts,
|
||||
images=images,
|
||||
tools=tools,
|
||||
structured_messages=structured_messages,
|
||||
shown_messages=guardrail_request.structured_messages,
|
||||
guardrail_response=guardrail_response,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import base64
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
|
|
@ -14,11 +15,13 @@ from litellm.integrations.custom_guardrail import (
|
|||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts, message_with_slot_texts
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -28,12 +31,36 @@ if TYPE_CHECKING:
|
|||
|
||||
_SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS: Final = 30.0
|
||||
_SANITIZE_FILE_QUEUED_STATUSES: Final = frozenset({"created", "in progress"})
|
||||
_PROTECT_ROLES: Final = frozenset({"system", "user", "assistant"})
|
||||
|
||||
|
||||
class PromptSecurityGuardrailMissingSecrets(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _inputs_with_structured_messages(
|
||||
inputs: GenericGuardrailAPIInputs, rewritten_messages: Sequence[AllMessageValues] | None
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if rewritten_messages is None:
|
||||
return inputs
|
||||
patched: Final[GenericGuardrailAPIInputs] = {
|
||||
**inputs,
|
||||
"structured_messages": list(rewritten_messages), # mutable-ok: the TypedDict field is declared as a list
|
||||
}
|
||||
return patched
|
||||
|
||||
|
||||
def _inputs_with_modifications(
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
modified_texts: list[str],
|
||||
rewritten_messages: Sequence[AllMessageValues] | None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if not modified_texts:
|
||||
return _inputs_with_structured_messages(inputs, rewritten_messages)
|
||||
with_texts: Final[GenericGuardrailAPIInputs] = {**inputs, "texts": modified_texts}
|
||||
return _inputs_with_structured_messages(with_texts, rewritten_messages)
|
||||
|
||||
|
||||
class _ProtectVerdict(TypedDict, total=False):
|
||||
"""One side (``prompt`` or ``response``) of an ``/api/protect`` verdict."""
|
||||
|
||||
|
|
@ -276,14 +303,39 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
detail="Blocked by Prompt Security, Violations: " + ", ".join(violations),
|
||||
)
|
||||
elif action == "modify":
|
||||
# Extract modified texts from modified_messages
|
||||
modified_messages: Final = result.get("modified_messages", [])
|
||||
modified_texts: Final = self._extract_texts_from_messages(modified_messages)
|
||||
if modified_texts:
|
||||
inputs["texts"] = modified_texts
|
||||
return _inputs_with_modifications(
|
||||
inputs,
|
||||
self._extract_texts_from_messages(modified_messages),
|
||||
self._structured_messages_with_modifications(structured_messages, modified_messages),
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
def _is_sent_to_protect(self, message: Mapping[str, object]) -> bool:
|
||||
return self.check_tool_results or message.get("role") in _PROTECT_ROLES
|
||||
|
||||
def _structured_messages_with_modifications(
|
||||
self,
|
||||
structured_messages: Sequence[AllMessageValues],
|
||||
modified_messages: Sequence[Mapping[str, object]],
|
||||
) -> tuple[AllMessageValues, ...] | None:
|
||||
sent_indices: Final = tuple(
|
||||
index for index, message in enumerate(structured_messages) if self._is_sent_to_protect(message)
|
||||
)
|
||||
if not sent_indices or len(sent_indices) != len(modified_messages):
|
||||
return None
|
||||
rewritten: Final = tuple(
|
||||
message_with_slot_texts(structured_messages[index], self._extract_texts_from_messages((modified,)))
|
||||
for index, modified in zip(sent_indices, modified_messages)
|
||||
)
|
||||
replacements: Final = MappingProxyType(
|
||||
{index: message for index, message in zip(sent_indices, rewritten) if message is not None}
|
||||
)
|
||||
if len(replacements) != len(sent_indices):
|
||||
return None
|
||||
return tuple(replacements.get(index, message) for index, message in enumerate(structured_messages))
|
||||
|
||||
async def _apply_guardrail_on_response(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
|
|
@ -347,19 +399,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
def _extract_texts_from_messages(self, messages: Sequence[Mapping[str, object]]) -> list[str]:
|
||||
"""Extract text content from messages."""
|
||||
texts: Final = []
|
||||
for message in messages:
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
texts.append(content)
|
||||
elif isinstance(content, list):
|
||||
for item in content:
|
||||
if isinstance(item, dict) and item.get("type") == "text":
|
||||
text = item.get("text")
|
||||
if text:
|
||||
texts.append(text)
|
||||
return texts
|
||||
return [text for message in messages for text in message_slot_texts(message)]
|
||||
|
||||
async def _process_standalone_images(self, images: list[str], user_api_key_alias: str | None) -> None:
|
||||
"""Process standalone images from inputs (data URLs)."""
|
||||
|
|
@ -681,14 +721,13 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
|
||||
This allows checking tool results for indirect prompt injection when enabled.
|
||||
"""
|
||||
supported_roles: Final = ["system", "user", "assistant"]
|
||||
filtered_messages: Final = []
|
||||
transformed_count = 0
|
||||
filtered_count = 0
|
||||
|
||||
for message in messages:
|
||||
role = message.get("role", "")
|
||||
if role in supported_roles:
|
||||
if role in _PROTECT_ROLES:
|
||||
filtered_messages.append(message)
|
||||
else:
|
||||
if self.check_tool_results:
|
||||
|
|
|
|||
|
|
@ -58,15 +58,6 @@ class UndeliverableStreamRewrite(Exception):
|
|||
self.guardrail_name: Final = guardrail_name
|
||||
|
||||
|
||||
class UnappliableRequestRewrite(Exception):
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
super().__init__(
|
||||
f"Guardrail '{guardrail_name}' rewrote the request in a way this endpoint cannot apply, "
|
||||
"so the request was rejected rather than sent unrewritten"
|
||||
)
|
||||
self.guardrail_name: Final = guardrail_name
|
||||
|
||||
|
||||
def _tool_call_shape(tool_call: object) -> tuple[object, object]:
|
||||
plain: Final = tool_call.model_dump() if isinstance(tool_call, BaseModel) else tool_call
|
||||
function: Final = plain.get("function") if isinstance(plain, Mapping) else None
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from typing import Any, Final, Literal
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Literal, cast # noqa: TID251 # JSON chat rows have no typed constructor across roles
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import TypedDict
|
||||
|
|
@ -158,12 +159,21 @@ def coerce_stream_holdback_value(value: Any) -> int:
|
|||
return 0
|
||||
|
||||
|
||||
def structured_messages_from_response(value: object) -> Sequence[AllMessageValues] | None:
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
if not all(isinstance(message, Mapping) and isinstance(message.get("role"), str) for message in value):
|
||||
return None
|
||||
return cast("Sequence[AllMessageValues]", value) # cast-ok: JSON rows checked for a role, the same trust texts get
|
||||
|
||||
|
||||
class GenericGuardrailAPIResponse:
|
||||
"""Response model for the Generic Guardrail API"""
|
||||
|
||||
texts: list[str] | None
|
||||
images: list[str] | None
|
||||
tools: list[GuardrailToolParam] | None
|
||||
structured_messages: Sequence[AllMessageValues] | None
|
||||
action: str
|
||||
blocked_reason: str | None
|
||||
stream_holdback_chars: list[int] | None
|
||||
|
|
@ -176,12 +186,14 @@ class GenericGuardrailAPIResponse:
|
|||
images: list[str] | None = None,
|
||||
tools: list[GuardrailToolParam] | None = None,
|
||||
stream_holdback_chars: list[int] | None = None,
|
||||
structured_messages: Sequence[AllMessageValues] | None = None,
|
||||
) -> None:
|
||||
self.action = action
|
||||
self.blocked_reason = blocked_reason
|
||||
self.texts = texts
|
||||
self.images = images
|
||||
self.tools = tools
|
||||
self.structured_messages = structured_messages
|
||||
# Number of trailing chars, indexed the same as ``texts``, that the
|
||||
# framework must withhold from streaming emission until the next
|
||||
# processing round (word-boundary safety for text transformations).
|
||||
|
|
@ -200,4 +212,5 @@ class GenericGuardrailAPIResponse:
|
|||
images=data.get("images"),
|
||||
tools=data.get("tools"),
|
||||
stream_holdback_chars=stream_holdback_chars,
|
||||
structured_messages=structured_messages_from_response(data.get("structured_messages")),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -283,8 +283,11 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
input_cost_per_video_token: float | None # for gemini omni models with video input
|
||||
input_cost_per_audio_per_second: float | None # only for vertex ai models
|
||||
input_cost_per_video_per_second: float | None # only for vertex ai models
|
||||
input_cost_per_audio_token_batches: ReadOnly[float | None]
|
||||
input_cost_per_image_token_batches: ReadOnly[float | None]
|
||||
input_cost_per_second: float | None # for OpenAI Speech models
|
||||
input_cost_per_token_batches: float | None
|
||||
input_cost_per_video_token_batches: ReadOnly[float | None]
|
||||
output_cost_per_token_batches: float | None
|
||||
output_cost_per_token: Required[float | None]
|
||||
output_cost_per_token_flex: float | None # OpenAI flex service tier pricing
|
||||
|
|
@ -3583,7 +3586,10 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
input_cost_per_video_per_second_above_128k_tokens: float | None = None
|
||||
input_cost_per_video_per_second_above_15s_interval: float | None = None
|
||||
input_cost_per_video_per_second_above_8s_interval: float | None = None
|
||||
input_cost_per_audio_token_batches: float | None = None
|
||||
input_cost_per_image_token_batches: float | None = None
|
||||
input_cost_per_token_batches: float | None = None
|
||||
input_cost_per_video_token_batches: float | None = None
|
||||
output_cost_per_token_batches: float | None = None
|
||||
output_cost_per_token_flex: float | None = None
|
||||
output_cost_per_token_priority: float | None = None
|
||||
|
|
|
|||
|
|
@ -5923,10 +5923,13 @@ def _get_model_info_helper(
|
|||
input_cost_per_audio_token=_model_info.get("input_cost_per_audio_token", None),
|
||||
input_cost_per_image_token=_model_info.get("input_cost_per_image_token", None),
|
||||
input_cost_per_video_token=_model_info.get("input_cost_per_video_token", None),
|
||||
input_cost_per_audio_token_batches=_model_info.get("input_cost_per_audio_token_batches", None),
|
||||
input_cost_per_image_token_batches=_model_info.get("input_cost_per_image_token_batches", None),
|
||||
input_cost_per_image=_model_info.get("input_cost_per_image", None),
|
||||
input_cost_per_audio_per_second=_model_info.get("input_cost_per_audio_per_second", None),
|
||||
input_cost_per_video_per_second=_model_info.get("input_cost_per_video_per_second", None),
|
||||
input_cost_per_token_batches=_model_info.get("input_cost_per_token_batches"),
|
||||
input_cost_per_video_token_batches=_model_info.get("input_cost_per_video_token_batches", None),
|
||||
output_cost_per_token_batches=_model_info.get("output_cost_per_token_batches"),
|
||||
output_cost_per_token=_output_cost_per_token,
|
||||
output_cost_per_token_flex=_model_info.get("output_cost_per_token_flex", None),
|
||||
|
|
|
|||
|
|
@ -25601,10 +25601,14 @@
|
|||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models"
|
||||
},
|
||||
"gemini-embedding-2-preview": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25615,13 +25619,14 @@
|
|||
"uses_embed_content": true
|
||||
},
|
||||
"gemini-embedding-2": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25633,10 +25638,14 @@
|
|||
"uses_embed_content": true
|
||||
},
|
||||
"vertex_ai/gemini-embedding-2-preview": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25648,13 +25657,14 @@
|
|||
"uses_embed_content": true
|
||||
},
|
||||
"vertex_ai/gemini-embedding-2": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25693,10 +25703,14 @@
|
|||
},
|
||||
"gemini/gemini-embedding-2-preview": {
|
||||
"deprecation_date": "2026-08-10",
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
@ -25709,10 +25723,14 @@
|
|||
"tpm": 10000000
|
||||
},
|
||||
"gemini/gemini-embedding-2": {
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"input_cost_per_audio_token": 6.5e-06,
|
||||
"input_cost_per_audio_token_batches": 3.25e-06,
|
||||
"input_cost_per_image_token": 4.5e-07,
|
||||
"input_cost_per_image_token_batches": 2.25e-07,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"input_cost_per_token_batches": 1e-07,
|
||||
"input_cost_per_video_token": 1.2e-05,
|
||||
"input_cost_per_video_token_batches": 6e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 8192,
|
||||
"max_tokens": 8192,
|
||||
|
|
|
|||
|
|
@ -249,6 +249,10 @@
|
|||
"type": "number",
|
||||
"minimum": 0
|
||||
},
|
||||
"input_cost_per_audio_token_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0
|
||||
},
|
||||
"input_cost_per_audio_token_priority": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
|
|
@ -276,6 +280,10 @@
|
|||
"type": "number",
|
||||
"minimum": 0
|
||||
},
|
||||
"input_cost_per_image_token_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0
|
||||
},
|
||||
"input_cost_per_pixel": {
|
||||
"type": "number",
|
||||
"minimum": 0
|
||||
|
|
@ -375,6 +383,14 @@
|
|||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"input_cost_per_video_token": {
|
||||
"type": "number",
|
||||
"minimum": 0
|
||||
},
|
||||
"input_cost_per_video_token_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0
|
||||
},
|
||||
"input_dbu_cost_per_token": {
|
||||
"type": "number",
|
||||
"minimum": 0
|
||||
|
|
|
|||
|
|
@ -81,6 +81,25 @@ def test_missing_execution_evidence_fails(tmp_path: Path, contents: str) -> None
|
|||
assert result.returncode == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("omitted_role", ("proxy_admin", "team_member", "internal_user_viewer"))
|
||||
def test_one_passing_management_case_cannot_hide_a_missing_actor(tmp_path: Path, omitted_role: str) -> None:
|
||||
suite: Final = ET.Element("testsuite")
|
||||
path: Final = "tests/e2e/management/test_jwt_management_e2e.py"
|
||||
case: Final = ET.SubElement(suite, "testcase", file=path)
|
||||
properties: Final = ET.SubElement(case, "properties")
|
||||
_ = ET.SubElement(
|
||||
properties,
|
||||
"property",
|
||||
name="management_node",
|
||||
value=f"{path}::TestJwtManagement::test_actor_subject_and_database_role[proxy_admin_viewer]",
|
||||
)
|
||||
report: Final = tmp_path / "report.xml"
|
||||
ET.ElementTree(suite).write(report)
|
||||
result: Final = subprocess.run([sys.executable, str(GATE), str(report), path], capture_output=True, text=True)
|
||||
assert result.returncode == 1
|
||||
assert f"test_actor_subject_and_database_role[{omitted_role}]" in result.stdout
|
||||
|
||||
|
||||
def test_short_values_are_written_without_masking_every_digit_in_the_log(tmp_path: Path) -> None:
|
||||
env_path: Final = tmp_path / ".env"
|
||||
|
||||
|
|
@ -141,6 +160,10 @@ def test_changed_suite_files_are_selected_unless_the_stack_cannot_run_them(
|
|||
(
|
||||
"tests/e2e/proxy_client.py",
|
||||
"tests/e2e/conftest.py",
|
||||
"tests/e2e/management/management_client.py",
|
||||
"tests/e2e/management/jwt_actors.py",
|
||||
"tests/e2e/management/conftest.py",
|
||||
"tests/e2e/coverage_registry/management_cases.py",
|
||||
"tests/e2e/pytest.ini",
|
||||
"tests/e2e/gateway/stage_mirror_ci_config.yml",
|
||||
".github/e2e-stack/up.sh",
|
||||
|
|
|
|||
329
tests/code_coverage_tests/test_provider_replay_harness.py
Normal file
329
tests/code_coverage_tests/test_provider_replay_harness.py
Normal file
|
|
@ -0,0 +1,329 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fixture_bundle import BundleRecorder, LoadedBundle, load_bundle, prepare_bundle
|
||||
from fixture_mode import current_test_key
|
||||
from fixture_profile import MatchProfile
|
||||
from provider_edge import REPLAY_MISS_STATUS, RecordEdge, ReplayEdge, ReplaySource
|
||||
from test_provider_edge import (
|
||||
CHAT_PATH,
|
||||
SSE_CHUNKS,
|
||||
STREAM_BODY,
|
||||
UPLOAD_PATH,
|
||||
call_edge,
|
||||
chunked_provider,
|
||||
fake_provider,
|
||||
json_object,
|
||||
provider_url,
|
||||
raw_stream_post,
|
||||
running_edge,
|
||||
this_tests_files,
|
||||
)
|
||||
|
||||
|
||||
class TestStrictIdentity:
|
||||
@pytest.mark.parametrize("path", [CHAT_PATH, "/anthropic/v1/messages"])
|
||||
def test_roundtrip_rejects_semantic_changes(self, tmp_path: Path, path: str) -> None:
|
||||
recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1")
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
original: Final = (
|
||||
b'{"model":"synthetic","messages":[{"role":"user",'
|
||||
b'"content":"2031-04-05 00000000-0000-0000-0000-000000000001"}],"options":[1,2]}'
|
||||
)
|
||||
headers: Final = {
|
||||
"content-type": "application/json",
|
||||
"accept": "application/json",
|
||||
"anthropic-version": "2023-06-01",
|
||||
"anthropic-beta": "feature-a",
|
||||
"openai-beta": "feature-b",
|
||||
"authorization": "Bearer synthetic-secret-one",
|
||||
}
|
||||
query: Final = "?part=one&part=two&blank="
|
||||
with fake_provider() as provider:
|
||||
mounts: Final = {"openai": provider_url(provider), "anthropic": provider_url(provider)}
|
||||
with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge:
|
||||
captured: Final = call_edge(edge, "POST", path + query, body=original, headers=headers)
|
||||
assert captured.status_code == 200
|
||||
assert json_object(captured.body)["echo"] == original.decode()
|
||||
loaded: Final = load_bundle(recorder.root, profile="stateless_v1")
|
||||
assert isinstance(loaded, LoadedBundle)
|
||||
assert loaded.manifest.match_profile == "stateless_v1"
|
||||
source: Final = ReplaySource(loaded)
|
||||
with running_edge(ReplayEdge(source), mounts) as edge:
|
||||
cases: Final = (
|
||||
(original.replace(b"2031-04-05", b"2032-06-07"), headers, query, "body"),
|
||||
(original.replace(b"000000000001", b"000000000002"), headers, query, "body"),
|
||||
(original.replace(b"[1,2]", b"[2,1]"), headers, query, "body"),
|
||||
(original.replace(b"synthetic", b"other"), headers, query, "body"),
|
||||
(original, headers, "?part=three&part=two&blank=", "query"),
|
||||
(original, headers, "?part=two&part=one&blank=", "query"),
|
||||
*(
|
||||
(original, {k: v for k, v in headers.items() if k != name}, query, "headers")
|
||||
for name in ("accept", "anthropic-version", "anthropic-beta", "openai-beta")
|
||||
),
|
||||
*(
|
||||
(original, {**headers, name: value}, query, "headers")
|
||||
for name in ("accept", "anthropic-version", "anthropic-beta", "openai-beta")
|
||||
for value in ("different", "")
|
||||
),
|
||||
(original, {**headers, "authorization": "Basic synthetic-secret-two"}, query, "auth"),
|
||||
(original, {k: v for k, v in headers.items() if k != "authorization"}, query, "auth"),
|
||||
)
|
||||
for rejected, reason in (
|
||||
(call_edge(edge, "POST", path + changed_query, body=body, headers=changed_headers), reason)
|
||||
for body, changed_headers, changed_query, reason in cases
|
||||
):
|
||||
assert rejected.status_code == REPLAY_MISS_STATUS
|
||||
assert reason in rejected.body.decode()
|
||||
assert b"synthetic-secret" not in rejected.body
|
||||
reordered: Final = json.dumps(dict(reversed(list(json_object(original).items())))).encode()
|
||||
accepted: Final = call_edge(
|
||||
edge, "POST", path + query, body=reordered, headers={k.upper(): v for k, v in headers.items()}
|
||||
)
|
||||
assert accepted.status_code == 200
|
||||
assert accepted.body == captured.body
|
||||
assert source.leftover_error(current_test_key()) is None
|
||||
assert len(provider.hits) == 1
|
||||
assert "synthetic-secret" not in "".join(file.read_text() for file in recorder.root.rglob("*.json"))
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
b'{"value":null}',
|
||||
b'{"value":""}',
|
||||
b'{"value":false}',
|
||||
b'{"value":0}',
|
||||
b'{"value":[]}',
|
||||
b'{"value":{}}',
|
||||
b'{"value":0.123456789012345678901}',
|
||||
b'{"value":0.123456789012345678902}',
|
||||
b'{"value":1e400}',
|
||||
b'{"value":1}',
|
||||
b'{"value":1e0}',
|
||||
b'{"value":-0}',
|
||||
b'{"value":1e9999999999999999999}',
|
||||
],
|
||||
)
|
||||
def test_json_values_remain_distinct(self, tmp_path: Path, body: bytes) -> None:
|
||||
recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1")
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
with fake_provider() as provider:
|
||||
mounts: Final = {"openai": provider_url(provider)}
|
||||
with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge:
|
||||
assert (
|
||||
call_edge(
|
||||
edge, "POST", CHAT_PATH, body=body, headers={"content-type": "application/json"}
|
||||
).status_code
|
||||
== 200
|
||||
)
|
||||
loaded: Final = load_bundle(recorder.root, profile="stateless_v1")
|
||||
assert isinstance(loaded, LoadedBundle)
|
||||
with running_edge(ReplayEdge(ReplaySource(loaded)), mounts) as edge:
|
||||
values: Final = (
|
||||
b"{}",
|
||||
b'{"value":null}',
|
||||
b'{"value":""}',
|
||||
b'{"value":false}',
|
||||
b'{"value":0}',
|
||||
b'{"value":[]}',
|
||||
b'{"value":{}}',
|
||||
b'{"value":0.123456789012345678901}',
|
||||
b'{"value":0.123456789012345678902}',
|
||||
b'{"value":1e400}',
|
||||
b'{"value":1}',
|
||||
b'{"value":1e0}',
|
||||
b'{"value":-0}',
|
||||
b'{"value":1e9999999999999999999}',
|
||||
)
|
||||
for rejected in (
|
||||
call_edge(edge, "POST", CHAT_PATH, body=value, headers={"content-type": "application/json"})
|
||||
for value in values
|
||||
if value != body
|
||||
):
|
||||
assert rejected.status_code == REPLAY_MISS_STATUS
|
||||
assert b"body" in rejected.body
|
||||
assert (
|
||||
call_edge(
|
||||
edge, "POST", CHAT_PATH, body=body, headers={"content-type": "application/json"}
|
||||
).status_code
|
||||
== 200
|
||||
)
|
||||
assert len(provider.hits) == 1
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"path,body,headers",
|
||||
[
|
||||
(UPLOAD_PATH, b"{}", {"content-type": "application/json"}),
|
||||
(CHAT_PATH + "?part=%FF", b"{}", {"content-type": "application/json"}),
|
||||
(CHAT_PATH + "?part=%FE", b"{}", {"content-type": "application/json"}),
|
||||
(CHAT_PATH, b"opaque", {"content-type": "application/octet-stream"}),
|
||||
(CHAT_PATH, b"--boundary", {"content-type": "multipart/form-data; boundary=boundary"}),
|
||||
(CHAT_PATH, b'{"x":1,"x":2}', {"content-type": "application/json"}),
|
||||
(CHAT_PATH, b"{}", {"content-type": "application/json", "x-custom-behavior": "synthetic-private-value"}),
|
||||
],
|
||||
)
|
||||
def test_ineligible_capture_never_calls_provider(
|
||||
self, tmp_path: Path, path: str, body: bytes, headers: dict[str, str]
|
||||
) -> None:
|
||||
recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1")
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
with fake_provider() as provider:
|
||||
with running_edge(RecordEdge(recorder, threading.Lock()), {"openai": provider_url(provider)}) as edge:
|
||||
result: Final = call_edge(edge, "POST", path, body=body, headers=headers)
|
||||
assert result.status_code == REPLAY_MISS_STATUS
|
||||
assert b"eligibility error" in result.body
|
||||
assert b"synthetic-private-value" not in result.body
|
||||
assert provider.hits == []
|
||||
assert this_tests_files(recorder.root) == []
|
||||
|
||||
def test_destination_is_part_of_actual_http_identity(self, tmp_path: Path) -> None:
|
||||
recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1")
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
with fake_provider() as provider:
|
||||
with running_edge(RecordEdge(recorder, threading.Lock()), {"openai": provider_url(provider)}) as edge:
|
||||
assert (
|
||||
call_edge(
|
||||
edge, "POST", CHAT_PATH, body=b"{}", headers={"content-type": "application/json"}
|
||||
).status_code
|
||||
== 200
|
||||
)
|
||||
loaded: Final = load_bundle(recorder.root, profile="stateless_v1")
|
||||
assert isinstance(loaded, LoadedBundle)
|
||||
with running_edge(ReplayEdge(ReplaySource(loaded)), {"openai": provider_url(provider) + "/other"}) as edge:
|
||||
result: Final = call_edge(
|
||||
edge, "POST", CHAT_PATH, body=b"{}", headers={"content-type": "application/json"}
|
||||
)
|
||||
assert result.status_code == REPLAY_MISS_STATUS
|
||||
assert b"upstream" in result.body
|
||||
assert len(provider.hits) == 1
|
||||
|
||||
def test_credentials_are_not_identity_and_fresh_process_replays(self, tmp_path: Path) -> None:
|
||||
recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1")
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
headers: Final = {
|
||||
"content-type": "application/json",
|
||||
"authorization": "bEaReR synthetic-token",
|
||||
"x-api-key": "synthetic-api-key",
|
||||
"cookie": "synthetic-cookie",
|
||||
}
|
||||
path: Final = CHAT_PATH + "?api_key=synthetic-query-secret&part=one&part=two"
|
||||
body: Final = b'{"model":"synthetic","messages":[]}'
|
||||
with fake_provider(echo_request=False) as provider:
|
||||
mounts: Final = {"openai": provider_url(provider)}
|
||||
with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge:
|
||||
captured: Final = call_edge(edge, "POST", path, body=body, headers=headers)
|
||||
assert captured.status_code == 200
|
||||
seen_headers, seen_body = provider.requests[0]
|
||||
assert {k.lower(): v for k, v in seen_headers.items()}.items() >= headers.items()
|
||||
assert seen_body == body
|
||||
assert provider.hits == ["POST " + path.removeprefix("/openai")]
|
||||
artifacts: Final = "".join(file.read_text() for file in recorder.root.rglob("*.json"))
|
||||
for secret in ("synthetic-token", "synthetic-api-key", "synthetic-cookie", "synthetic-query-secret"):
|
||||
assert secret not in artifacts
|
||||
child: Final = subprocess.run(
|
||||
[
|
||||
sys.executable,
|
||||
"-c",
|
||||
"""
|
||||
import json, sys
|
||||
from pathlib import Path
|
||||
from fixture_bundle import LoadedBundle, load_bundle
|
||||
from provider_edge import ProviderRequestObservation, observed_provider_edge, replay_leftover_error
|
||||
from test_provider_edge import call_edge
|
||||
from fixture_profile import MatchProfile
|
||||
from fixture_mode import current_test_key
|
||||
loaded = load_bundle(Path(sys.argv[1]), profile="stateless_v1")
|
||||
assert isinstance(loaded, LoadedBundle)
|
||||
with observed_provider_edge(ProviderRequestObservation("synthetic"), mode_raw="replay", bundle_dir=Path(sys.argv[1]), bind_host="127.0.0.1", advertise_host="127.0.0.1", mounts={"openai": sys.argv[2]}) as edge:
|
||||
response = call_edge(edge, "POST", sys.argv[3], body=sys.argv[4].encode(), headers=json.loads(sys.argv[5]))
|
||||
assert response.status_code == 200
|
||||
print(response.body.decode())
|
||||
assert replay_leftover_error(mode_raw="replay", bundle_dir=Path(sys.argv[1]), test_key=current_test_key()) is None
|
||||
""",
|
||||
str(recorder.root),
|
||||
provider_url(provider),
|
||||
path.replace("synthetic-query-secret", "new-query-credential"),
|
||||
body.decode(),
|
||||
json.dumps({**headers, "authorization": "Bearer another-credential", "x-api-key": "another-key"}),
|
||||
],
|
||||
env={
|
||||
**os.environ,
|
||||
"PYTHONPATH": str(Path(__file__).resolve().parents[1] / "e2e"),
|
||||
"E2E_REPLAY_MATCH_PROFILE": "stateless_v1",
|
||||
},
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
assert child.returncode == 0, child.stderr
|
||||
assert child.stdout.strip().encode() == captured.body
|
||||
assert len(provider.hits) == 1
|
||||
|
||||
@pytest.mark.parametrize("profile,other", [("legacy", "stateless_v1"), ("stateless_v1", "legacy")])
|
||||
def test_profiles_cannot_load_each_others_bundles(
|
||||
self, tmp_path: Path, profile: MatchProfile, other: MatchProfile
|
||||
) -> None:
|
||||
from fixture_bundle import UnreadableBundle
|
||||
|
||||
recorder: Final = prepare_bundle(tmp_path / profile, profile=profile)
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
mismatch: Final = load_bundle(recorder.root, profile=other)
|
||||
assert isinstance(mismatch, UnreadableBundle)
|
||||
assert "profile mismatch" in mismatch.reason
|
||||
assert "re-record" in mismatch.reason
|
||||
|
||||
@pytest.mark.parametrize("abort_after", [None, 2])
|
||||
def test_strict_stream_preserves_chunks_and_truncation(self, tmp_path: Path, abort_after: int | None) -> None:
|
||||
recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1")
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
with chunked_provider(abort_after=abort_after) as provider:
|
||||
mounts: Final = {"anthropic": provider_url(provider)}
|
||||
with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge:
|
||||
_, captured, captured_ending = raw_stream_post(edge.port, "/anthropic/v1/messages", STREAM_BODY)
|
||||
loaded: Final = load_bundle(recorder.root, profile="stateless_v1")
|
||||
assert isinstance(loaded, LoadedBundle)
|
||||
source: Final = ReplaySource(loaded)
|
||||
with running_edge(ReplayEdge(source), mounts) as edge:
|
||||
_, replayed, ending = raw_stream_post(edge.port, "/anthropic/v1/messages", STREAM_BODY)
|
||||
assert captured == replayed == list(SSE_CHUNKS[:abort_after])
|
||||
assert ending == captured_ending
|
||||
assert (ending == "terminated") == (abort_after is None)
|
||||
assert source.leftover_error(current_test_key()) is None
|
||||
assert len(provider.hits) == 1
|
||||
|
||||
def test_auth_scheme_survives_missing_credentials(self, tmp_path: Path) -> None:
|
||||
recorder: Final = prepare_bundle(tmp_path / "strict", profile="stateless_v1")
|
||||
assert isinstance(recorder, BundleRecorder)
|
||||
headers: Final = {"content-type": "application/json", "authorization": "Bearer"}
|
||||
with fake_provider() as provider:
|
||||
mounts: Final = {"openai": provider_url(provider)}
|
||||
with running_edge(RecordEdge(recorder, threading.Lock()), mounts) as edge:
|
||||
assert call_edge(edge, "POST", CHAT_PATH, body=b"{}", headers=headers).status_code == 200
|
||||
loaded: Final = load_bundle(recorder.root, profile="stateless_v1")
|
||||
assert isinstance(loaded, LoadedBundle)
|
||||
with running_edge(ReplayEdge(ReplaySource(loaded)), mounts) as edge:
|
||||
for result in (
|
||||
call_edge(edge, "POST", CHAT_PATH, body=b"{}", headers={**headers, "authorization": scheme})
|
||||
for scheme in ("Basic", "Digest")
|
||||
):
|
||||
assert result.status_code == REPLAY_MISS_STATUS
|
||||
assert b"auth" in result.body
|
||||
assert (
|
||||
call_edge(
|
||||
edge,
|
||||
"POST",
|
||||
CHAT_PATH,
|
||||
body=b"{}",
|
||||
headers={**headers, "authorization": "bEaReR synthetic-token"},
|
||||
).status_code
|
||||
== 200
|
||||
)
|
||||
assert len(provider.hits) == 1
|
||||
|
|
@ -60,7 +60,11 @@ The suites run against a live proxy, so bring one up first by running the litell
|
|||
|
||||
Keycloak's password grant is a test-only provisioning shortcut, not a production login recommendation. The `litellm-e2e-admin` client adds the proxy's admin scope; the normal client does not. Never reuse this permissive realm outside an isolated test stack.
|
||||
|
||||
Management tests can use the shared `idp` and `jwt_identity` fixtures. Each test gets a unique Keycloak group/user and a matching proxy user/team. Setup and fallback cleanup use the master key; the operations and read-backs being tested must explicitly use `caller_key=idp.access_token(jwt_identity, client_id=ADMIN_CLIENT_ID)` (or a member token). See `management/test_jwt_management_e2e.py` for create/read/update/clear/delete and tenant-denial examples. A group claim alone is not database team membership: permission tests explicitly add the member and prove an allowed read before asserting the denied write.
|
||||
Management tests can bind a credential once with `client.with_caller(Caller(...))`; direct calls, delegated helpers and replica read-backs then retain that caller. Explicit `caller_key` arguments override the binding. Keep the original master-backed client for bootstrap and cleanup. `actor_factory` lazily provisions database roles and tenant memberships, with `database_role` tokens carrying no groups and `group_scoped` actors retaining the existing team route gate. Token minting is explicit through `actor.mint_caller(idp)`. The factory runs requests without backend retries and reports cleanup failures. `coverage_registry/management_cases.py` records exact canary nodes and non-secret actor labels; the CI execution assertion rejects a missing or skipped actor row
|
||||
|
||||
For the opt-in browser profile, start the existing IdP first, then run `.github/e2e-stack/oidc-profile.sh "$PROXY_BASE_URL" <server-command>`. The wrapper creates a confidential client with an exact `/sso/callback` redirect and S256 PKCE, passes the client secret only through the child process environment, and removes the client on exit. It uses the existing generic OIDC handler with `GENERIC_USER_ID_ATTRIBUTE=sub`. Preserve the IdP's PostgreSQL data across restarts
|
||||
|
||||
`tests/e2e/ui/playwright.oidc.config.ts` uses an already running OIDC stack and separate storage/output files. Supply `E2E_OIDC_UI_URL`, `JWT_ISSUER`, `E2E_OIDC_USERNAME` and `E2E_OIDC_PASSWORD` for a seeded actor. Its setup follows the real login and callback path. The current Python canary qualifies browser-client configuration and token/userinfo identity mapping; browser journey specs under `ui/oidc/` are a separate coverage step
|
||||
|
||||
Every successful IdP create immediately registers cleanup, including partial setup failures. Cleanup failures emit warnings. Tokens are minted on demand, and the expiration test waits relative to the token's actual `exp` with a bounded clock-drift check. To check first-attempt behavior locally, run both files with `--reruns 0`:
|
||||
|
||||
|
|
@ -232,3 +236,15 @@ Before you push
|
|||
4. Capture screenshots of the test run and attach them to the PR as proof
|
||||
|
||||
5. If a test fails because it surfaced a real issue in the product, flag that explicitly in the PR rather than reworking the test until it passes
|
||||
|
||||
### Strict stateless replay matching
|
||||
|
||||
Set `E2E_REPLAY_MATCH_PROFILE=stateless_v1` for both recording and replay to bind OpenAI `/v1/chat/completions` and Anthropic `/v1/messages` requests to their upstream destination, ordered query pairs, semantic headers and literal JSON content. The default remains `legacy`. Strict bundles use format 5 and cannot load as legacy bundles; select the matching profile or re-record with `E2E_FIXTURE_MODE=record`. Missing profile metadata never enrolls a legacy bundle in strict matching
|
||||
|
||||
Strict matching preserves dates, UUIDs, hashes, model names, tool arguments, array order and omitted/null/empty/false/zero values. JSON object key order and header name casing may change. The strict body uses tagged JSON values so number precision and JSON types survive persistence, including exact numeric spelling and numbers larger than a floating-point value. Invalid UTF-8 query values fail eligibility. Duplicate JSON keys, unsupported endpoints, non-JSON bodies and unknown semantic headers fail eligibility before contacting a provider
|
||||
|
||||
The semantic header set is `content-type`, `accept`, `anthropic-version`, `anthropic-beta` and `openai-beta`, including missing versus present values. Authorization records presence and the case-insensitive scheme; `x-api-key` records presence only. Credential values and cookies are excluded. Credential query values are redacted while their position and field name remain in the identity. Never use real customer inputs in fixture qualification
|
||||
|
||||
Excluded transport and telemetry headers are `host`, `content-length`, `connection`, `accept-encoding`, `user-agent`, `traceparent`, `tracestate`, `x-request-id`, `x-client-request-id` and `x-stainless-*`. Inbound transfer-encoding is unsupported; send JSON with content-length framing. The destination represents host identity and the relay carries original body bytes. Replay does not verify credentials, SDK timeout/retry behavior, transport performance, model availability or stateful remote IDs. Live relay uses original request bytes and header values, never the stored identity
|
||||
|
||||
Strict replay harness regression tests live in `tests/code_coverage_tests/test_provider_replay_harness.py`. The CircleCI `provider_replay_harness` job runs them alongside the existing legacy harness files with `--noconftest -o pythonpath=tests/e2e`; they need only synthetic HTTP providers and temporary fixture storage
|
||||
|
|
|
|||
151
tests/e2e/coverage_registry/management_cases.py
Normal file
151
tests/e2e/coverage_registry/management_cases.py
Normal file
|
|
@ -0,0 +1,151 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Final, Literal
|
||||
|
||||
CredentialKind = Literal["master", "idp_admin", "direct_jwt", "virtual_key", "dashboard_session"]
|
||||
DependencyProfile = Literal["management_only", "real_oidc_browser", "external_provider_required"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ManagementCase:
|
||||
node: str
|
||||
credential_kind: CredentialKind
|
||||
actor: str
|
||||
profile: str
|
||||
method: Literal["GET", "POST"]
|
||||
path: str
|
||||
operation_family: str
|
||||
dependency_profile: DependencyProfile = "management_only"
|
||||
|
||||
|
||||
JWT_FILE: Final = "tests/e2e/management/test_jwt_management_e2e.py"
|
||||
JWT_CLASS: Final = f"{JWT_FILE}::TestJwtManagement"
|
||||
ACTORS: Final = (
|
||||
"proxy_admin",
|
||||
"proxy_admin_viewer",
|
||||
"organization_admin",
|
||||
"team_admin",
|
||||
"team_member",
|
||||
"internal_user",
|
||||
"internal_user_viewer",
|
||||
"unrelated_user",
|
||||
)
|
||||
MANAGEMENT_CASES: Final = tuple(
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_actor_subject_and_database_role[{role}]",
|
||||
credential_kind="direct_jwt",
|
||||
actor=role,
|
||||
profile="database_role",
|
||||
method="GET",
|
||||
path="/user/info",
|
||||
operation_family="identity",
|
||||
)
|
||||
for role in ACTORS
|
||||
) + (
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_admin_viewer_reads_but_cannot_update",
|
||||
credential_kind="direct_jwt",
|
||||
actor="proxy_admin_viewer",
|
||||
profile="database_role",
|
||||
method="POST",
|
||||
path="/key/update",
|
||||
operation_family="denial",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_admin_creates_reads_updates_clears_and_deletes_a_key[direct_jwt]",
|
||||
credential_kind="direct_jwt",
|
||||
actor="proxy_admin",
|
||||
profile="group_scoped",
|
||||
method="POST",
|
||||
path="/key/generate",
|
||||
operation_family="key_lifecycle",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_admin_creates_reads_updates_clears_and_deletes_a_key[virtual_key]",
|
||||
credential_kind="virtual_key",
|
||||
actor="proxy_admin",
|
||||
profile="group_scoped",
|
||||
method="POST",
|
||||
path="/key/generate",
|
||||
operation_family="key_lifecycle",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_two_actor_sets_keep_tenants_and_keys_isolated",
|
||||
credential_kind="direct_jwt",
|
||||
actor="team_member",
|
||||
profile="group_scoped",
|
||||
method="GET",
|
||||
path="/key/info",
|
||||
operation_family="tenant_isolation",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_member_cannot_write_and_another_team_cannot_read_the_key",
|
||||
credential_kind="direct_jwt",
|
||||
actor="team_member",
|
||||
profile="group_scoped",
|
||||
method="POST",
|
||||
path="/key/update",
|
||||
operation_family="tenant_isolation",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_multi_group_actor_keeps_exact_memberships",
|
||||
credential_kind="master",
|
||||
actor="bootstrap",
|
||||
profile="group_scoped",
|
||||
method="GET",
|
||||
path="/team/info",
|
||||
operation_family="memberships",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_successful_actor_cleanup_removes_owned_state",
|
||||
credential_kind="master",
|
||||
actor="bootstrap",
|
||||
profile="failure_cleanup",
|
||||
method="GET",
|
||||
path="/team/info",
|
||||
operation_family="cleanup",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_partial_setup_removes_previously_created_identities[group]",
|
||||
credential_kind="idp_admin",
|
||||
actor="idp_admin",
|
||||
profile="failure_cleanup",
|
||||
method="POST",
|
||||
path="/groups",
|
||||
operation_family="cleanup",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_partial_setup_removes_previously_created_identities[user]",
|
||||
credential_kind="idp_admin",
|
||||
actor="idp_admin",
|
||||
profile="failure_cleanup",
|
||||
method="POST",
|
||||
path="/users",
|
||||
operation_family="cleanup",
|
||||
),
|
||||
ManagementCase(
|
||||
node=f"{JWT_CLASS}::test_oidc_browser_profile_identity_mapping",
|
||||
credential_kind="direct_jwt",
|
||||
actor="internal_user",
|
||||
profile="oidc_configuration",
|
||||
method="GET",
|
||||
path="/protocol/openid-connect/userinfo",
|
||||
operation_family="oidc_identity",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def canonical_node(node: str) -> str:
|
||||
return node if node.startswith("tests/e2e/") else f"tests/e2e/{node}"
|
||||
|
||||
|
||||
def case_properties(node: str) -> tuple[tuple[str, str], ...]:
|
||||
case: Final = next((case for case in MANAGEMENT_CASES if case.node == canonical_node(node)), None)
|
||||
if case is None:
|
||||
return ()
|
||||
return (
|
||||
("management_node", case.node),
|
||||
("credential_kind", case.credential_kind),
|
||||
("actor", case.actor),
|
||||
("auth_profile", case.profile),
|
||||
("dependency_profile", case.dependency_profile),
|
||||
)
|
||||
|
|
@ -90,3 +90,11 @@
|
|||
- {id: mgmt.mcp_toolset.update.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3098", rationale: "Narrowing the tools to one entry reads back exactly that entry"}
|
||||
- {id: mgmt.mcp_toolset.update.clear_persists, module: mgmt, tier: P0, surface: api, assertions: [clear_persists], source: "mcp_management_endpoints.py:3098", fail_before_fix: proven, rationale: "An explicit null clears the stored description; the update used to drop null and keep the old value"}
|
||||
- {id: mgmt.mcp_toolset.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "mcp_management_endpoints.py:3149", rationale: "A deleted toolset is gone by id and from the list on every replica"}
|
||||
|
||||
- {id: mgmt.user.jwt.database_roles, module: mgmt, tier: P0, surface: api, assertions: [database_roles], source: "auth/handle_jwt.py", rationale: "User-only JWT subjects retain their seeded database roles and memberships"}
|
||||
- {id: mgmt.key.jwt.viewer_denied, module: mgmt, tier: P0, surface: api, assertions: [viewer_denied], source: "auth/route_checks.py", rationale: "An admin viewer can read a key but cannot update it or change stored state"}
|
||||
- {id: mgmt.user.oidc.identity_mapping, module: mgmt, tier: P0, surface: api, assertions: [identity_mapping], source: "tests/e2e/idp.py", rationale: "IdP configuration canary only: confidential-client token and userinfo subjects match the seeded user; application SSO is separate"}
|
||||
- {id: mgmt.team.jwt.tenant_isolation, module: mgmt, tier: P0, surface: api, assertions: [tenant_isolation], source: "auth/handle_jwt.py", rationale: "Isolated team actors read their own key and receive 403 for the other tenant key"}
|
||||
- {id: mgmt.team.jwt.multiple_memberships, module: mgmt, tier: P0, surface: api, assertions: [multiple_memberships], source: "auth/handle_jwt.py", rationale: "A multi-group actor has exactly the configured memberships without admin scope"}
|
||||
- {id: mgmt.user.jwt.cleanup, module: mgmt, tier: P0, surface: api, assertions: [cleanup], source: "management_endpoints/internal_user_endpoints.py", rationale: "Owned users teams organizations keys and IdP objects disappear after successful cleanup"}
|
||||
- {id: mgmt.user.jwt.partial_cleanup, module: mgmt, tier: P0, surface: api, assertions: [partial_cleanup], source: "auth/handle_jwt.py", rationale: "Partial identity setup removes the group and user created before failure"}
|
||||
|
|
|
|||
|
|
@ -16,9 +16,11 @@ requests itself imports.
|
|||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from collections.abc import Callable, Generator, Iterator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Generator, Generic, Iterator, Literal, NewType, Protocol, TypeVar, cast
|
||||
from typing import Final, Generic, Literal, NewType, Protocol, TypeVar, cast
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
|
@ -36,8 +38,8 @@ class Headers(BaseModel):
|
|||
|
||||
class AuthHeaders(Headers):
|
||||
# litellm accepts either; set whichever the call needs, leave the other None.
|
||||
authorization: str | None = None
|
||||
x_litellm_api_key: str | None = Field(default=None, alias="x-litellm-api-key")
|
||||
authorization: str | None = Field(default=None, repr=False)
|
||||
x_litellm_api_key: str | None = Field(default=None, alias="x-litellm-api-key", repr=False)
|
||||
|
||||
|
||||
class AnthropicHeaders(AuthHeaders):
|
||||
|
|
@ -292,6 +294,22 @@ def _params(params: BaseModel | None) -> dict[str, str]:
|
|||
|
||||
TRANSIENT_STATUSES: frozenset[int] = frozenset({529})
|
||||
RETRY_ATTEMPTS: int = 3
|
||||
_QUALIFICATION: Final[ContextVar[bool]] = ContextVar("e2e_qualification", default=False)
|
||||
|
||||
|
||||
def retry_attempts(default: int) -> int:
|
||||
return 1 if _QUALIFICATION.get() else default
|
||||
|
||||
|
||||
@contextmanager
|
||||
def without_retries() -> Generator[None]:
|
||||
token: Final = _QUALIFICATION.set(True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_QUALIFICATION.reset(token)
|
||||
|
||||
|
||||
RETRY_BACKOFF_SECONDS: float = 0.5
|
||||
|
||||
|
||||
|
|
@ -319,7 +337,7 @@ def request_with_retry[T: RetryableResponse](
|
|||
hang should surface as a hang instead of doubling the wall clock. Every
|
||||
retry prints, so flakiness stays visible in the run log instead of
|
||||
vanishing into green."""
|
||||
for attempt in range(1, RETRY_ATTEMPTS):
|
||||
for attempt in range(1, retry_attempts(RETRY_ATTEMPTS)):
|
||||
resp = issue()
|
||||
if resp.status_code not in TRANSIENT_STATUSES:
|
||||
return resp
|
||||
|
|
@ -414,6 +432,7 @@ def get_external[R: BaseModel](
|
|||
url: str,
|
||||
*,
|
||||
response_type: type[R],
|
||||
headers: BaseModel | None = None,
|
||||
timeout: float = 30.0,
|
||||
) -> Result[R]:
|
||||
"""GET an absolute URL outside the proxy (e.g. a public /.well-known document).
|
||||
|
|
@ -422,7 +441,7 @@ def get_external[R: BaseModel](
|
|||
try:
|
||||
resp = requests.get(
|
||||
url,
|
||||
headers={"Accept": "application/json"},
|
||||
headers={"Accept": "application/json", **(_headers(headers) if headers is not None else {})},
|
||||
timeout=timeout,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
|
|
|
|||
|
|
@ -30,9 +30,11 @@ from datetime import datetime, timedelta, timezone
|
|||
from pathlib import Path
|
||||
from typing import Annotated, Final, Literal
|
||||
|
||||
from fixture_profile import MatchProfile, StrictIdentity
|
||||
from pydantic import BaseModel, Field, JsonValue
|
||||
|
||||
BUNDLE_FORMAT_VERSION: Final = 4
|
||||
STRICT_BUNDLE_FORMAT_VERSION: Final = 5
|
||||
MAX_BUNDLE_AGE: Final = timedelta(days=7)
|
||||
MANIFEST_FILENAME: Final = "manifest.json"
|
||||
|
||||
|
|
@ -41,6 +43,7 @@ class Manifest(BaseModel):
|
|||
format_version: int
|
||||
recorded_at: datetime
|
||||
harness_version: str
|
||||
match_profile: MatchProfile = "legacy"
|
||||
|
||||
|
||||
class RecordedRequest(BaseModel):
|
||||
|
|
@ -69,6 +72,7 @@ class RecordedRequest(BaseModel):
|
|||
file_name: str | None = None
|
||||
file_sha256: str | None = None
|
||||
file_bytes: int | None = None
|
||||
strict_identity: StrictIdentity | None = None
|
||||
|
||||
|
||||
class RecordedHttpResponse(BaseModel):
|
||||
|
|
@ -100,9 +104,7 @@ class RecordedStreamedResponse(BaseModel):
|
|||
truncated: str | None = None
|
||||
|
||||
|
||||
type RecordedResponse = Annotated[
|
||||
RecordedHttpResponse | RecordedStreamedResponse, Field(discriminator="kind")
|
||||
]
|
||||
type RecordedResponse = Annotated[RecordedHttpResponse | RecordedStreamedResponse, Field(discriminator="kind")]
|
||||
|
||||
|
||||
class Interaction(BaseModel):
|
||||
|
|
@ -152,6 +154,7 @@ class BundleRecorder:
|
|||
manifest, so record mode never reads (or merges into) an existing bundle."""
|
||||
|
||||
root: Path
|
||||
profile: MatchProfile = "legacy"
|
||||
_ordinals: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
def record(self, *, test_key: str, request: RecordedRequest, response: RecordedResponse) -> None:
|
||||
|
|
@ -162,7 +165,12 @@ class BundleRecorder:
|
|||
directory.mkdir(parents=True, exist_ok=True)
|
||||
interaction = Interaction(request=request, response=response)
|
||||
target = directory / interaction_filename(ordinal, request)
|
||||
target.write_text(interaction.model_dump_json(indent=2), encoding="utf-8")
|
||||
target.write_text(
|
||||
interaction.model_dump_json(
|
||||
indent=2, exclude={"request": {"strict_identity"}} if self.profile == "legacy" else None
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -171,7 +179,7 @@ class UnsafeBundleDir:
|
|||
reason: str
|
||||
|
||||
|
||||
def prepare_bundle(root: Path) -> BundleRecorder | UnsafeBundleDir:
|
||||
def prepare_bundle(root: Path, *, profile: MatchProfile = "legacy") -> BundleRecorder | UnsafeBundleDir:
|
||||
"""Start a fresh bundle at ``root`` for record mode: wipe whatever bundle is
|
||||
there and write a new manifest. Refuses to wipe a directory that is neither
|
||||
empty nor a bundle (no manifest.json), so a mistyped E2E_FIXTURE_DIR can
|
||||
|
|
@ -188,12 +196,15 @@ def prepare_bundle(root: Path) -> BundleRecorder | UnsafeBundleDir:
|
|||
shutil.rmtree(root)
|
||||
root.mkdir(parents=True)
|
||||
manifest = Manifest(
|
||||
format_version=BUNDLE_FORMAT_VERSION,
|
||||
format_version=BUNDLE_FORMAT_VERSION if profile == "legacy" else STRICT_BUNDLE_FORMAT_VERSION,
|
||||
match_profile=profile,
|
||||
recorded_at=datetime.now(timezone.utc),
|
||||
harness_version=harness_version(),
|
||||
)
|
||||
(root / MANIFEST_FILENAME).write_text(manifest.model_dump_json(indent=2), encoding="utf-8")
|
||||
return BundleRecorder(root=root)
|
||||
(root / MANIFEST_FILENAME).write_text(
|
||||
manifest.model_dump_json(indent=2, exclude={"match_profile"} if profile == "legacy" else None), encoding="utf-8"
|
||||
)
|
||||
return BundleRecorder(root=root, profile=profile)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -226,25 +237,30 @@ def _read_manifest(root: Path) -> Manifest | UnreadableBundle:
|
|||
return UnreadableBundle(reason=f"{MANIFEST_FILENAME} is invalid: {exc}")
|
||||
|
||||
|
||||
def _supported_manifest(root: Path) -> Manifest | UnreadableBundle:
|
||||
def _supported_manifest(root: Path, profile: MatchProfile = "legacy") -> Manifest | UnreadableBundle:
|
||||
"""The manifest, refused when it was written under a different format version.
|
||||
A bundle is atomic (record wipes and rewrites the whole directory and never
|
||||
merges), so a foreign version is a hard reject rather than a partial read."""
|
||||
manifest = _read_manifest(root)
|
||||
if isinstance(manifest, UnreadableBundle):
|
||||
return manifest
|
||||
if manifest.format_version != BUNDLE_FORMAT_VERSION:
|
||||
expected_version: Final = BUNDLE_FORMAT_VERSION if profile == "legacy" else STRICT_BUNDLE_FORMAT_VERSION
|
||||
if manifest.match_profile != profile:
|
||||
return UnreadableBundle(
|
||||
reason="match profile mismatch; select the recorded E2E_REPLAY_MATCH_PROFILE or re-record"
|
||||
)
|
||||
if manifest.format_version != expected_version:
|
||||
return UnreadableBundle(
|
||||
reason=(
|
||||
f"format_version {manifest.format_version} != supported {BUNDLE_FORMAT_VERSION}; "
|
||||
f"format_version {manifest.format_version} != supported {expected_version}; "
|
||||
"re-record with E2E_FIXTURE_MODE=record"
|
||||
)
|
||||
)
|
||||
return manifest
|
||||
|
||||
|
||||
def check_freshness(root: Path, *, now: datetime) -> BundleFreshness:
|
||||
manifest = _supported_manifest(root)
|
||||
def check_freshness(root: Path, *, now: datetime, profile: MatchProfile = "legacy") -> BundleFreshness:
|
||||
manifest = _supported_manifest(root, profile)
|
||||
if isinstance(manifest, UnreadableBundle):
|
||||
return manifest
|
||||
recorded_at = (
|
||||
|
|
@ -269,16 +285,27 @@ class LoadedBundle:
|
|||
interactions: dict[str, tuple[Interaction, ...]]
|
||||
|
||||
|
||||
def load_bundle(root: Path) -> LoadedBundle | UnreadableBundle:
|
||||
manifest = _supported_manifest(root)
|
||||
def load_bundle(root: Path, *, profile: MatchProfile = "legacy") -> LoadedBundle | UnreadableBundle:
|
||||
manifest = _supported_manifest(root, profile)
|
||||
if isinstance(manifest, UnreadableBundle):
|
||||
return manifest
|
||||
interactions = {
|
||||
directory.name: tuple(
|
||||
Interaction.model_validate_json(file.read_text(encoding="utf-8"))
|
||||
for file in sorted(directory.glob("*.json"))
|
||||
)
|
||||
for directory in sorted(root.iterdir())
|
||||
if directory.is_dir()
|
||||
}
|
||||
try:
|
||||
interactions = {
|
||||
directory.name: tuple(
|
||||
Interaction.model_validate_json(file.read_text(encoding="utf-8"))
|
||||
for file in sorted(directory.glob("*.json"))
|
||||
)
|
||||
for directory in sorted(root.iterdir())
|
||||
if directory.is_dir()
|
||||
}
|
||||
except (ValueError, OSError):
|
||||
if profile == "legacy":
|
||||
raise
|
||||
return UnreadableBundle(reason="invalid stateless_v1 interaction; re-record with the selected profile")
|
||||
if any(
|
||||
(item.request.strict_identity is not None) != (profile == "stateless_v1")
|
||||
for items in interactions.values()
|
||||
for item in items
|
||||
):
|
||||
return UnreadableBundle(reason="request identity/profile mismatch; re-record with the selected profile")
|
||||
return LoadedBundle(manifest=manifest, interactions=interactions)
|
||||
|
|
|
|||
|
|
@ -23,9 +23,8 @@ from dataclasses import dataclass
|
|||
from functools import reduce
|
||||
from typing import Final
|
||||
|
||||
from pydantic import JsonValue
|
||||
|
||||
from fixture_bundle import RecordedRequest
|
||||
from pydantic import JsonValue
|
||||
|
||||
VOLATILE_HEADER_NAMES: Final[frozenset[str]] = frozenset(
|
||||
{
|
||||
|
|
@ -123,6 +122,12 @@ class CanonicalRequest:
|
|||
|
||||
|
||||
def canonicalize(request: RecordedRequest) -> CanonicalRequest:
|
||||
if request.strict_identity is not None:
|
||||
return CanonicalRequest(
|
||||
method=request.method,
|
||||
path=request.path,
|
||||
content=json.dumps(request.strict_identity.model_dump(mode="json"), sort_keys=True, separators=(",", ":")),
|
||||
)
|
||||
file_identity: Final[JsonValue | None] = (
|
||||
None
|
||||
if request.file_name is None and request.file_sha256 is None
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ from fixture_bundle import (
|
|||
check_freshness,
|
||||
format_age,
|
||||
)
|
||||
from fixture_profile import match_profile
|
||||
|
||||
type FixtureMode = Literal["live", "record", "replay"]
|
||||
|
||||
|
|
@ -82,6 +83,7 @@ def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datet
|
|||
Called at collection time (conftest pytest_sessionstart) so a stale or missing
|
||||
bundle fails the whole run up front, naming the bundle age, instead of failing
|
||||
every test individually."""
|
||||
match_profile()
|
||||
mode = parse_fixture_mode(mode_raw)
|
||||
match mode:
|
||||
case InvalidFixtureMode(value=value):
|
||||
|
|
@ -89,7 +91,7 @@ def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datet
|
|||
case "live" | "record":
|
||||
return None
|
||||
case "replay":
|
||||
freshness = check_freshness(bundle_dir, now=now)
|
||||
freshness = check_freshness(bundle_dir, now=now, profile=match_profile())
|
||||
match freshness:
|
||||
case FreshBundle():
|
||||
return None
|
||||
|
|
@ -110,6 +112,7 @@ def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datet
|
|||
def fixture_report_lines(mode_raw: str, bundle_dir: Path, *, now: datetime) -> list[str]:
|
||||
"""pytest report-header lines; empty in live mode so an unset
|
||||
E2E_FIXTURE_MODE keeps today's output byte-identical."""
|
||||
match_profile()
|
||||
mode = parse_fixture_mode(mode_raw)
|
||||
match mode:
|
||||
case InvalidFixtureMode() | "live":
|
||||
|
|
@ -117,7 +120,7 @@ def fixture_report_lines(mode_raw: str, bundle_dir: Path, *, now: datetime) -> l
|
|||
case "record":
|
||||
return [f"e2e fixture mode: record -> {bundle_dir}"]
|
||||
case "replay":
|
||||
freshness = check_freshness(bundle_dir, now=now)
|
||||
freshness = check_freshness(bundle_dir, now=now, profile=match_profile())
|
||||
match freshness:
|
||||
case FreshBundle(manifest=manifest):
|
||||
return [
|
||||
|
|
|
|||
179
tests/e2e/fixture_profile.py
Normal file
179
tests/e2e/fixture_profile.py
Normal file
|
|
@ -0,0 +1,179 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import parse_qsl, urlsplit
|
||||
|
||||
from pydantic import BaseModel, JsonValue, TypeAdapter
|
||||
|
||||
type MatchProfile = Literal["legacy", "stateless_v1"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NumberToken:
|
||||
literal: str
|
||||
|
||||
|
||||
type ExactJson = dict[str, ExactJson] | list[ExactJson] | str | bool | NumberToken | None
|
||||
|
||||
SEMANTIC_HEADERS: Final = frozenset({"content-type", "accept", "anthropic-version", "anthropic-beta", "openai-beta"})
|
||||
AUTH_HEADERS: Final = frozenset({"authorization", "x-api-key"})
|
||||
EXCLUDED_HEADERS: Final = frozenset(
|
||||
{
|
||||
"host",
|
||||
"content-length",
|
||||
"transfer-encoding",
|
||||
"connection",
|
||||
"accept-encoding",
|
||||
"user-agent",
|
||||
"traceparent",
|
||||
"tracestate",
|
||||
"x-request-id",
|
||||
"x-client-request-id",
|
||||
"cookie",
|
||||
}
|
||||
)
|
||||
CREDENTIAL_QUERY: Final = frozenset(
|
||||
{
|
||||
"api_key",
|
||||
"api-key",
|
||||
"apikey",
|
||||
"key",
|
||||
"token",
|
||||
"access_token",
|
||||
"signature",
|
||||
"password",
|
||||
"secret",
|
||||
"credentials",
|
||||
"authorization",
|
||||
"sig",
|
||||
"client_secret",
|
||||
"aws_access_key_id",
|
||||
"aws_secret_access_key",
|
||||
"aws_session_token",
|
||||
}
|
||||
)
|
||||
JSON_VALUE: Final[TypeAdapter[ExactJson]] = TypeAdapter(ExactJson)
|
||||
|
||||
|
||||
def match_profile() -> MatchProfile:
|
||||
raw: Final = os.environ.get("E2E_REPLAY_MATCH_PROFILE", "legacy")
|
||||
if raw in ("legacy", "stateless_v1"):
|
||||
return raw
|
||||
raise ValueError("E2E_REPLAY_MATCH_PROFILE must be legacy or stateless_v1")
|
||||
|
||||
|
||||
class StrictIdentity(BaseModel):
|
||||
upstream: str
|
||||
mount: str
|
||||
query: tuple[tuple[str, str], ...]
|
||||
headers: dict[str, str]
|
||||
auth: dict[str, str]
|
||||
body_present: bool
|
||||
body: JsonValue
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class IneligibleRequest:
|
||||
reason: str
|
||||
|
||||
|
||||
def _unique_object(pairs: list[tuple[str, ExactJson]]) -> dict[str, ExactJson]:
|
||||
if len({key for key, _ in pairs}) != len(pairs):
|
||||
raise ValueError("duplicate JSON object keys")
|
||||
return dict(pairs)
|
||||
|
||||
|
||||
def _invalid_constant(value: str) -> ExactJson:
|
||||
raise ValueError("nonfinite JSON number")
|
||||
|
||||
|
||||
def _exact_value(value: ExactJson) -> JsonValue:
|
||||
match value:
|
||||
case dict():
|
||||
return {"object": {key: _exact_value(item) for key, item in value.items()}}
|
||||
case list():
|
||||
return {"array": [_exact_value(item) for item in value]}
|
||||
case bool():
|
||||
return {"boolean": value}
|
||||
case NumberToken(literal=literal):
|
||||
return {"number": literal}
|
||||
case str():
|
||||
return {"string": value}
|
||||
case None:
|
||||
return None
|
||||
|
||||
|
||||
def strict_identity(
|
||||
*,
|
||||
method: str,
|
||||
path: str,
|
||||
query: str,
|
||||
headers: Mapping[str, str],
|
||||
body: bytes | None,
|
||||
mount: str,
|
||||
upstream_base: str,
|
||||
) -> StrictIdentity | IneligibleRequest:
|
||||
if (mount, path, method.upper()) not in {
|
||||
("openai", "/openai/v1/chat/completions", "POST"),
|
||||
("anthropic", "/anthropic/v1/messages", "POST"),
|
||||
}:
|
||||
return IneligibleRequest("unsupported endpoint or method")
|
||||
lowered: Final = {key.lower(): value for key, value in headers.items()}
|
||||
if len(lowered) != len(headers):
|
||||
return IneligibleRequest("duplicate header names")
|
||||
if any(
|
||||
key not in SEMANTIC_HEADERS | AUTH_HEADERS | EXCLUDED_HEADERS and not key.startswith("x-stainless-")
|
||||
for key in lowered
|
||||
):
|
||||
return IneligibleRequest("unsupported semantic header")
|
||||
if "transfer-encoding" in lowered:
|
||||
return IneligibleRequest("unsupported request transfer-encoding; send a content-length framed JSON body")
|
||||
authorization: Final = lowered.get("authorization")
|
||||
if authorization is not None and authorization.partition(" ")[0].lower() not in {"bearer", "basic", "digest"}:
|
||||
return IneligibleRequest("unsupported authorization scheme")
|
||||
destination: Final = urlsplit(upstream_base)
|
||||
if destination.username or destination.password or destination.query or destination.fragment:
|
||||
return IneligibleRequest("upstream destination contains credentials, query or fragment")
|
||||
if destination.scheme not in ("http", "https") or not destination.netloc:
|
||||
return IneligibleRequest("unsupported upstream destination")
|
||||
if body and lowered.get("content-type", "").split(";", 1)[0].strip().lower() != "application/json":
|
||||
return IneligibleRequest("unsupported body content-type; stateless_v1 requires JSON")
|
||||
try:
|
||||
parsed: Final = (
|
||||
JSON_VALUE.validate_python(
|
||||
json.loads(
|
||||
body,
|
||||
object_pairs_hook=_unique_object,
|
||||
parse_constant=_invalid_constant,
|
||||
parse_float=NumberToken,
|
||||
parse_int=NumberToken,
|
||||
)
|
||||
)
|
||||
if body
|
||||
else None
|
||||
)
|
||||
except (ValueError, UnicodeError):
|
||||
return IneligibleRequest("invalid JSON or duplicate JSON object keys")
|
||||
if body and not isinstance(parsed, dict):
|
||||
return IneligibleRequest("stateless inference requires a JSON object")
|
||||
try:
|
||||
query_pairs: Final = tuple(parse_qsl(query, keep_blank_values=True, errors="strict"))
|
||||
except UnicodeError:
|
||||
return IneligibleRequest("invalid UTF-8 query encoding")
|
||||
return StrictIdentity(
|
||||
upstream=upstream_base,
|
||||
mount=mount,
|
||||
query=tuple((key, "<credential>" if key.lower() in CREDENTIAL_QUERY else value) for key, value in query_pairs),
|
||||
headers={key: value for key, value in lowered.items() if key in SEMANTIC_HEADERS},
|
||||
auth={
|
||||
key: (value.partition(" ")[0].lower() if key == "authorization" else "present")
|
||||
for key, value in lowered.items()
|
||||
if key in AUTH_HEADERS
|
||||
},
|
||||
body_present=bool(body),
|
||||
body=_exact_value(parsed),
|
||||
)
|
||||
279
tests/e2e/idp.py
279
tests/e2e/idp.py
|
|
@ -2,11 +2,18 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import os
|
||||
import secrets
|
||||
import signal
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import warnings
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from contextlib import ExitStack
|
||||
from dataclasses import dataclass, field, replace
|
||||
from types import FrameType
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
|
|
@ -14,11 +21,15 @@ from e2e_http import (
|
|||
AuthHeaders,
|
||||
ExternalWrite,
|
||||
NetworkError,
|
||||
NoBody,
|
||||
Result,
|
||||
Success,
|
||||
UnknownApiError,
|
||||
delete_external,
|
||||
get_external,
|
||||
post_form_external,
|
||||
post_json_external,
|
||||
unwrap,
|
||||
)
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
|
@ -46,7 +57,9 @@ class TokenGrantForm(BaseModel):
|
|||
grant_type: Literal["password"] = "password"
|
||||
client_id: str
|
||||
username: str
|
||||
password: str
|
||||
password: str = Field(repr=False)
|
||||
client_secret: str | None = Field(default=None, repr=False)
|
||||
scope: str | None = None
|
||||
|
||||
|
||||
class TokenResponse(BaseModel):
|
||||
|
|
@ -63,7 +76,7 @@ class GroupCreateBody(BaseModel):
|
|||
|
||||
class PasswordCredential(BaseModel):
|
||||
type: Literal["password"] = "password"
|
||||
value: str
|
||||
value: str = Field(repr=False)
|
||||
temporary: bool = False
|
||||
|
||||
|
||||
|
|
@ -101,8 +114,20 @@ class Identity:
|
|||
user_id: str
|
||||
username: str
|
||||
password: str = field(repr=False)
|
||||
group: str
|
||||
group_id: str
|
||||
groups: tuple[str, ...]
|
||||
group_ids: tuple[str, ...]
|
||||
|
||||
@property
|
||||
def group(self) -> str:
|
||||
if len(self.groups) != 1:
|
||||
raise ValueError("A single-group identity is required")
|
||||
return self.groups[0]
|
||||
|
||||
@property
|
||||
def group_id(self) -> str:
|
||||
if len(self.group_ids) != 1:
|
||||
raise ValueError("A single-group identity is required")
|
||||
return self.group_ids[0]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -111,6 +136,10 @@ class Keycloak:
|
|||
realm: str
|
||||
admin_username: str
|
||||
admin_password: str = field(repr=False)
|
||||
strict_cleanup: bool = False
|
||||
|
||||
def with_strict_cleanup(self) -> Keycloak:
|
||||
return replace(self, strict_cleanup=True)
|
||||
|
||||
@property
|
||||
def issuer(self) -> str:
|
||||
|
|
@ -150,7 +179,9 @@ class Keycloak:
|
|||
f"group {name}",
|
||||
)
|
||||
|
||||
def create_user(self, *, username: str, email: str, password: str, group: str) -> str:
|
||||
def create_user(
|
||||
self, *, username: str, email: str, password: str, group: str | None = None, groups: tuple[str, ...] = ()
|
||||
) -> str:
|
||||
return created_id(
|
||||
post_json_external(
|
||||
self._admin_url("/users"),
|
||||
|
|
@ -158,7 +189,7 @@ class Keycloak:
|
|||
json=UserCreateBody(
|
||||
username=username,
|
||||
email=email,
|
||||
groups=(group,),
|
||||
groups=(group,) if group is not None else groups,
|
||||
credentials=(PasswordCredential(value=password),),
|
||||
),
|
||||
),
|
||||
|
|
@ -171,14 +202,28 @@ class Keycloak:
|
|||
def delete_group(self, group_id: str) -> None:
|
||||
self._delete(f"/groups/{group_id}")
|
||||
|
||||
def assert_absent(self, kind: Literal["users", "groups", "clients"], resource_id: str) -> None:
|
||||
result: Final = get_external(
|
||||
self._admin_url(f"/{kind}/{resource_id}"),
|
||||
headers=self._admin_headers(),
|
||||
response_type=NoBody,
|
||||
)
|
||||
assert isinstance(result, UnknownApiError) and result.status_code == 404, (
|
||||
f"Owned IdP {kind} still exists: {result}"
|
||||
)
|
||||
|
||||
def _delete(self, path: str) -> None:
|
||||
try:
|
||||
headers: Final = self._admin_headers()
|
||||
except pytest.fail.Exception as exc:
|
||||
if self.strict_cleanup:
|
||||
raise RuntimeError(f"Keycloak cleanup could not authenticate for {path}") from exc
|
||||
warnings.warn(f"Keycloak cleanup could not authenticate for {path}: {exc}", RuntimeWarning, stacklevel=2)
|
||||
return
|
||||
result: Final = delete_external(self._admin_url(path), headers=headers)
|
||||
if result.status_code not in (204, 404):
|
||||
if self.strict_cleanup:
|
||||
raise RuntimeError(f"Keycloak cleanup failed for {path}: HTTP {result.status_code}")
|
||||
warnings.warn(
|
||||
f"Keycloak cleanup failed for {path}: HTTP {result.status_code} {result.body[:300]}",
|
||||
RuntimeWarning,
|
||||
|
|
@ -188,15 +233,34 @@ class Keycloak:
|
|||
def provision(self, *, marker: str, group: str, defer: Callable[[Callable[[], object]], None]) -> Identity:
|
||||
"""Create `group` and a user in it, credentialed with a password generated
|
||||
for this test alone, and hand back the identity a token can be minted for."""
|
||||
group_id: Final = self.create_group(group)
|
||||
defer(lambda: self.delete_group(group_id))
|
||||
return self.provision_groups(marker=marker, groups=(group,), defer=defer)
|
||||
|
||||
def provision_groups(
|
||||
self, *, marker: str, groups: tuple[str, ...], defer: Callable[[Callable[[], object]], None]
|
||||
) -> Identity:
|
||||
def provision_group(name: str) -> str:
|
||||
created: Final = self.create_group(name)
|
||||
defer(lambda: self.delete_group(created))
|
||||
return created
|
||||
|
||||
group_ids: Final = tuple(provision_group(group) for group in groups)
|
||||
return self.provision_user(marker=marker, groups=groups, group_ids=group_ids, defer=defer)
|
||||
|
||||
def provision_user(
|
||||
self,
|
||||
*,
|
||||
marker: str,
|
||||
groups: tuple[str, ...],
|
||||
group_ids: tuple[str, ...],
|
||||
defer: Callable[[Callable[[], object]], None],
|
||||
) -> Identity:
|
||||
username: Final = f"e2e-jwt-user-{marker}"
|
||||
password: Final = secrets.token_urlsafe(24)
|
||||
user_id: Final = self.create_user(
|
||||
username=username, email=f"{username}@example.com", password=password, group=group
|
||||
username=username, email=f"{username}@example.com", password=password, groups=groups
|
||||
)
|
||||
defer(lambda: self.delete_user(user_id))
|
||||
return Identity(user_id=user_id, username=username, password=password, group=group, group_id=group_id)
|
||||
return Identity(user_id=user_id, username=username, password=password, groups=groups, group_ids=group_ids)
|
||||
|
||||
def access_token(
|
||||
self, identity: Identity, *, client_id: str = TESTS_CLIENT_ID, issuer_host: str | None = None
|
||||
|
|
@ -211,6 +275,65 @@ class Keycloak:
|
|||
)
|
||||
return self._token(result, f"a token for {identity.username}")
|
||||
|
||||
def discovery(self) -> Discovery:
|
||||
return unwrap(get_external(f"{self.issuer}/.well-known/openid-configuration", response_type=Discovery))
|
||||
|
||||
def browser_client(self, *, callback_url: str, defer: Callable[[Callable[[], object]], None]) -> BrowserClient:
|
||||
client: Final = BrowserClient(
|
||||
client_id=f"e2e-browser-{secrets.token_hex(8)}",
|
||||
secret=secrets.token_urlsafe(32),
|
||||
callback_url=callback_url,
|
||||
)
|
||||
resource_id: Final = created_id(
|
||||
post_json_external(
|
||||
self._admin_url("/clients"),
|
||||
headers=self._admin_headers(),
|
||||
json=BrowserClientBody(
|
||||
clientId=client.client_id,
|
||||
secret=client.secret,
|
||||
redirectUris=(callback_url,),
|
||||
),
|
||||
),
|
||||
"browser client",
|
||||
)
|
||||
defer(lambda: self._delete(f"/clients/{resource_id}"))
|
||||
configured: Final = unwrap(
|
||||
get_external(
|
||||
self._admin_url(f"/clients/{resource_id}"),
|
||||
headers=self._admin_headers(),
|
||||
response_type=BrowserClientBody,
|
||||
)
|
||||
)
|
||||
assert configured.redirect_uris == (callback_url,)
|
||||
assert configured.standard_flow_enabled and not configured.public_client
|
||||
assert configured.attributes.pkce == "S256"
|
||||
return client
|
||||
|
||||
def browser_token(self, identity: Identity, client: BrowserClient) -> str:
|
||||
return self._token(
|
||||
post_form_external(
|
||||
self.token_url(self.realm),
|
||||
form=TokenGrantForm(
|
||||
client_id=client.client_id,
|
||||
client_secret=client.secret,
|
||||
username=identity.username,
|
||||
password=identity.password,
|
||||
scope="openid email",
|
||||
),
|
||||
response_type=TokenResponse,
|
||||
),
|
||||
"browser-profile identity mapping",
|
||||
)
|
||||
|
||||
def userinfo(self, token: str) -> UserInfo:
|
||||
return unwrap(
|
||||
get_external(
|
||||
f"{self.issuer}/protocol/openid-connect/userinfo",
|
||||
headers=AuthHeaders(authorization=f"Bearer {token}"),
|
||||
response_type=UserInfo,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def keycloak_from_env() -> Keycloak:
|
||||
admin_username: Final = os.environ.get(KEYCLOAK_ADMIN_USER_ENV, "").strip()
|
||||
|
|
@ -226,3 +349,137 @@ def keycloak_from_env() -> Keycloak:
|
|||
admin_username=admin_username,
|
||||
admin_password=admin_password,
|
||||
)
|
||||
|
||||
|
||||
class TokenClaims(BaseModel):
|
||||
sub: str
|
||||
iss: str
|
||||
aud: str | tuple[str, ...]
|
||||
exp: int
|
||||
scope: str = ""
|
||||
groups: tuple[str, ...] = ()
|
||||
|
||||
|
||||
class Discovery(BaseModel):
|
||||
issuer: str
|
||||
authorization_endpoint: str
|
||||
token_endpoint: str
|
||||
userinfo_endpoint: str
|
||||
jwks_uri: str
|
||||
|
||||
|
||||
class UserInfo(BaseModel):
|
||||
sub: str
|
||||
email: str
|
||||
|
||||
|
||||
class BrowserAttributes(BaseModel):
|
||||
pkce: str = Field(default="S256", alias="pkce.code.challenge.method")
|
||||
|
||||
|
||||
class AudienceConfig(BaseModel):
|
||||
audience: str = Field(default="litellm-e2e", alias="included.custom.audience")
|
||||
access_token: str = Field(default="true", alias="access.token.claim")
|
||||
id_token: str = Field(default="false", alias="id.token.claim")
|
||||
|
||||
|
||||
class AudienceMapper(BaseModel):
|
||||
name: str = "litellm-audience"
|
||||
protocol: str = "openid-connect"
|
||||
mapper: str = Field(default="oidc-audience-mapper", alias="protocolMapper")
|
||||
config: AudienceConfig = Field(default_factory=AudienceConfig)
|
||||
|
||||
|
||||
class BrowserClientBody(BaseModel):
|
||||
client_id: str = Field(alias="clientId")
|
||||
secret: str = Field(repr=False)
|
||||
redirect_uris: tuple[str, ...] = Field(alias="redirectUris")
|
||||
enabled: bool = True
|
||||
public_client: bool = Field(default=False, alias="publicClient")
|
||||
standard_flow_enabled: bool = Field(default=True, alias="standardFlowEnabled")
|
||||
direct_access_grants_enabled: bool = Field(default=True, alias="directAccessGrantsEnabled")
|
||||
default_client_scopes: tuple[str, ...] = Field(default=("email", "basic"), alias="defaultClientScopes")
|
||||
attributes: BrowserAttributes = Field(default_factory=BrowserAttributes)
|
||||
protocol_mappers: tuple[AudienceMapper, ...] = Field(default=(AudienceMapper(),), alias="protocolMappers")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BrowserClient:
|
||||
client_id: str
|
||||
secret: str = field(repr=False)
|
||||
callback_url: str
|
||||
|
||||
def environment(self, discovery: Discovery) -> dict[str, str]:
|
||||
return {
|
||||
"GENERIC_CLIENT_ID": self.client_id,
|
||||
"GENERIC_CLIENT_SECRET": self.secret,
|
||||
"GENERIC_USER_ID_ATTRIBUTE": "sub",
|
||||
"GENERIC_AUTHORIZATION_ENDPOINT": discovery.authorization_endpoint,
|
||||
"GENERIC_TOKEN_ENDPOINT": discovery.token_endpoint,
|
||||
"GENERIC_USERINFO_ENDPOINT": discovery.userinfo_endpoint,
|
||||
"GENERIC_CLIENT_USE_PKCE": "true",
|
||||
"GENERIC_SCOPE": "openid email",
|
||||
}
|
||||
|
||||
|
||||
def token_claims(token: str) -> TokenClaims:
|
||||
payload: Final = token.split(".")[1]
|
||||
return TokenClaims.model_validate_json(base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4)))
|
||||
|
||||
|
||||
def _signal_process_group(process_id: int, signum: int) -> bool:
|
||||
try:
|
||||
os.killpg(process_id, signum)
|
||||
except ProcessLookupError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _stop_process_group(child: subprocess.Popen[bytes]) -> None:
|
||||
_signal_process_group(child.pid, signal.SIGTERM)
|
||||
deadline: Final = time.monotonic() + 5
|
||||
while _process_group_exists(child.pid):
|
||||
child.poll()
|
||||
if time.monotonic() >= deadline:
|
||||
_signal_process_group(child.pid, signal.SIGKILL)
|
||||
break
|
||||
time.sleep(0.05)
|
||||
child.wait()
|
||||
|
||||
|
||||
def _process_group_exists(process_id: int) -> bool:
|
||||
try:
|
||||
os.killpg(process_id, 0)
|
||||
except ProcessLookupError:
|
||||
return False
|
||||
except PermissionError:
|
||||
return True
|
||||
return True
|
||||
|
||||
|
||||
def run_oidc_profile(proxy_url: str, command: list[str]) -> int:
|
||||
idp: Final = keycloak_from_env().with_strict_cleanup()
|
||||
with ExitStack() as cleanup:
|
||||
|
||||
def terminate(signum: int, frame: FrameType | None) -> None:
|
||||
raise SystemExit(128 + signum)
|
||||
|
||||
previous: Final = signal.signal(signal.SIGTERM, terminate)
|
||||
cleanup.callback(signal.signal, signal.SIGTERM, previous)
|
||||
|
||||
def defer(callback: Callable[[], object]) -> None:
|
||||
cleanup.callback(callback)
|
||||
|
||||
client: Final = idp.browser_client(callback_url=f"{proxy_url.rstrip('/')}/sso/callback", defer=defer)
|
||||
environment: Final = {**os.environ, **client.environment(idp.discovery()), "PROXY_BASE_URL": proxy_url}
|
||||
with subprocess.Popen(command, env=environment, start_new_session=True) as child:
|
||||
try:
|
||||
return child.wait()
|
||||
finally:
|
||||
_stop_process_group(child)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if len(sys.argv) < 3:
|
||||
raise SystemExit("Usage: idp.py PROXY_URL COMMAND [ARG ...]; requires a running test IdP")
|
||||
raise SystemExit(run_oidc_profile(sys.argv[1], sys.argv[2:]))
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from __future__ import annotations
|
|||
from collections.abc import Iterable
|
||||
|
||||
import pytest
|
||||
from coverage_registry.management_cases import case_properties
|
||||
|
||||
# Hardcoded because the runner image copies tests/e2e/ to /app/e2e, so nothing
|
||||
# at runtime names this suite's place in the repo. test_junit_properties.py
|
||||
|
|
@ -94,7 +95,7 @@ def result_properties(item: pytest.Item) -> tuple[tuple[str, str], ...]:
|
|||
("package", package_from_nodeid(item.nodeid)),
|
||||
("covers", ",".join(covers_from_item(item))),
|
||||
("source", source_from_item(item)),
|
||||
)
|
||||
) + case_properties(item.nodeid)
|
||||
|
||||
|
||||
def attach_result_properties(item: pytest.Item) -> None:
|
||||
|
|
|
|||
|
|
@ -5,8 +5,14 @@ holds the shared ProxyClient so `resources` / `scoped_key` clean up keys, teams,
|
|||
users, and orgs this suite creates.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from collections.abc import Generator
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from e2e_http import without_retries
|
||||
from idp import Keycloak
|
||||
from lifecycle import ResourceManager
|
||||
from management.jwt_actors import ActorFactory
|
||||
from management_client import ManagementClient, build_client
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
|
|
@ -21,3 +27,14 @@ def pytest_configure(config: pytest.Config) -> None:
|
|||
@pytest.fixture(scope="session")
|
||||
def client(proxy: ProxyClient) -> ManagementClient:
|
||||
return build_client(proxy)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def actor_factory(proxy: ProxyClient, idp: Keycloak) -> Generator[ActorFactory]:
|
||||
bootstrap: Final = build_client(proxy)
|
||||
resources: Final = ResourceManager(client=proxy, strict_cleanup=True)
|
||||
with without_retries():
|
||||
try:
|
||||
yield ActorFactory(bootstrap=bootstrap, idp=idp, resources=resources)
|
||||
finally:
|
||||
resources.teardown()
|
||||
|
|
|
|||
175
tests/e2e/management/jwt_actors.py
Normal file
175
tests/e2e/management/jwt_actors.py
Normal file
|
|
@ -0,0 +1,175 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Literal
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import NoBody, unwrap
|
||||
from idp import ADMIN_CLIENT_ID, TESTS_CLIENT_ID, Identity, Keycloak
|
||||
from lifecycle import ResourceManager
|
||||
from management.management_client import ManagementClient
|
||||
from models import (
|
||||
KeyGenerateBody,
|
||||
KeyGenerateResponse,
|
||||
OrgDeleteBody,
|
||||
OrgDeleteResponse,
|
||||
OrgMemberAddBody,
|
||||
OrgMemberEntry,
|
||||
OrgNewBody,
|
||||
TeamDeleteBody,
|
||||
TeamMemberAddBody,
|
||||
TeamMemberEntry,
|
||||
TeamNewBody,
|
||||
UserNewBody,
|
||||
UserRole,
|
||||
)
|
||||
from proxy_client import Caller
|
||||
|
||||
ActorRole = Literal[
|
||||
"proxy_admin",
|
||||
"proxy_admin_viewer",
|
||||
"organization_admin",
|
||||
"team_admin",
|
||||
"team_member",
|
||||
"internal_user",
|
||||
"internal_user_viewer",
|
||||
"unrelated_user",
|
||||
]
|
||||
ActorProfile = Literal["database_role", "group_scoped"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Tenant:
|
||||
organization_id: str
|
||||
team_id: str
|
||||
group_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Actor:
|
||||
identity: Identity
|
||||
role: ActorRole
|
||||
global_role: UserRole
|
||||
profile: ActorProfile
|
||||
tenants: tuple[Tenant, ...]
|
||||
|
||||
def mint_caller(self, idp: Keycloak) -> Caller:
|
||||
return Caller(
|
||||
credential=idp.access_token(
|
||||
self.identity, client_id=ADMIN_CLIENT_ID if self.role == "proxy_admin" else TESTS_CLIENT_ID
|
||||
),
|
||||
kind="direct_jwt",
|
||||
role=self.role,
|
||||
tenant=self.tenants[0].organization_id if self.tenants else None,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ActorFactory:
|
||||
bootstrap: ManagementClient
|
||||
idp: Keycloak
|
||||
resources: ResourceManager
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.bootstrap.proxy.caller is not None:
|
||||
raise ValueError("Actor bootstrap requires a separately held master client")
|
||||
|
||||
def key(self, tenant: Tenant | None = None, *, user_id: str | None = None) -> KeyGenerateResponse:
|
||||
created: Final = unwrap(
|
||||
self.bootstrap.generate_key(
|
||||
KeyGenerateBody(
|
||||
team_id=tenant.team_id if tenant is not None else None,
|
||||
user_id=user_id,
|
||||
key_alias=f"e2e-actor-key-{unique_marker()}",
|
||||
)
|
||||
)
|
||||
)
|
||||
self.resources.defer(lambda: self.bootstrap.delete_key_strict(created.key, missing_ok=True))
|
||||
return created
|
||||
|
||||
def tenant(self) -> Tenant:
|
||||
marker: Final = unique_marker()
|
||||
organization_id: Final = self.bootstrap.create_org(OrgNewBody(organization_alias=f"e2e-organization-{marker}"))
|
||||
self.resources.defer(
|
||||
lambda: unwrap(
|
||||
self.bootstrap.proxy.transport.delete(
|
||||
"/organization/delete",
|
||||
headers=self.bootstrap.proxy.management_headers(),
|
||||
json=OrgDeleteBody(organization_ids=[organization_id]),
|
||||
response_type=OrgDeleteResponse,
|
||||
)
|
||||
)
|
||||
)
|
||||
team_id: Final = self.bootstrap.proxy.create_team(
|
||||
TeamNewBody(team_alias=f"e2e-team-{marker}", organization_id=organization_id)
|
||||
)
|
||||
self.resources.defer(
|
||||
lambda: unwrap(
|
||||
self.bootstrap.proxy.transport.post(
|
||||
"/team/delete",
|
||||
headers=self.bootstrap.proxy.management_headers(),
|
||||
json=TeamDeleteBody(team_ids=[team_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
)
|
||||
self.bootstrap.delete_team_member(team_id, self.bootstrap.user_info().user_id)
|
||||
group_id: Final = self.idp.create_group(team_id)
|
||||
self.resources.defer(lambda: self.idp.with_strict_cleanup().delete_group(group_id))
|
||||
return Tenant(organization_id=organization_id, team_id=team_id, group_id=group_id)
|
||||
|
||||
def create(
|
||||
self, role: ActorRole, *, tenants: tuple[Tenant, ...] = (), profile: ActorProfile = "database_role"
|
||||
) -> Actor:
|
||||
if role in ("team_admin", "team_member", "organization_admin") and not tenants:
|
||||
raise ValueError("A membership actor requires a tenant")
|
||||
identity: Final = self.idp.with_strict_cleanup().provision_user(
|
||||
marker=unique_marker(),
|
||||
groups=tuple(tenant.team_id for tenant in tenants) if profile == "group_scoped" else (),
|
||||
group_ids=tuple(tenant.group_id for tenant in tenants) if profile == "group_scoped" else (),
|
||||
defer=self.resources.defer,
|
||||
)
|
||||
global_role: Final[UserRole] = (
|
||||
role
|
||||
if role in ("proxy_admin", "proxy_admin_viewer", "internal_user", "internal_user_viewer")
|
||||
else "internal_user"
|
||||
)
|
||||
self.bootstrap.create_user(
|
||||
UserNewBody(
|
||||
user_id=identity.user_id,
|
||||
user_email=f"{identity.username}@example.com",
|
||||
user_role=global_role,
|
||||
auto_create_key=False,
|
||||
)
|
||||
)
|
||||
self.resources.defer(lambda: self.bootstrap.delete_user_strict(identity.user_id))
|
||||
for tenant in tenants:
|
||||
unwrap(
|
||||
self.bootstrap.proxy.transport.post(
|
||||
"/organization/member_add",
|
||||
headers=self.bootstrap.proxy.management_headers(),
|
||||
json=OrgMemberAddBody(
|
||||
organization_id=tenant.organization_id,
|
||||
member=OrgMemberEntry(
|
||||
user_id=identity.user_id,
|
||||
role="org_admin" if role == "organization_admin" else "internal_user",
|
||||
),
|
||||
),
|
||||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
unwrap(
|
||||
self.bootstrap.proxy.transport.post(
|
||||
"/team/member_add",
|
||||
headers=self.bootstrap.proxy.management_headers(),
|
||||
json=TeamMemberAddBody(
|
||||
team_id=tenant.team_id,
|
||||
member=TeamMemberEntry(
|
||||
user_id=identity.user_id,
|
||||
role="admin" if role == "team_admin" else "user",
|
||||
),
|
||||
),
|
||||
response_type=NoBody,
|
||||
)
|
||||
)
|
||||
return Actor(identity=identity, role=role, global_role=global_role, profile=profile, tenants=tenants)
|
||||
|
|
@ -7,7 +7,8 @@ llm-only key hitting a management route).
|
|||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
import warnings
|
||||
from dataclasses import dataclass, field, replace
|
||||
|
||||
import jwt
|
||||
from e2e_config import MASTER_KEY
|
||||
|
|
@ -20,6 +21,7 @@ from e2e_http import (
|
|||
StreamingResponse,
|
||||
Success,
|
||||
UnknownApiError,
|
||||
retry_attempts,
|
||||
unwrap,
|
||||
)
|
||||
from models import (
|
||||
|
|
@ -81,7 +83,7 @@ from models import (
|
|||
UserNewResponse,
|
||||
UserUpdateBody,
|
||||
)
|
||||
from proxy_client import ProxyClient
|
||||
from proxy_client import Caller, ProxyClient
|
||||
|
||||
MODEL_ACCESS_DENIED_MARKER = "key_model_access_denied"
|
||||
ROUTE_NOT_ALLOWED_MARKER = "not allowed to call this route"
|
||||
|
|
@ -98,7 +100,7 @@ class DashboardSession:
|
|||
its bearer on every subsequent call, the claims it renders the signed-in user
|
||||
from, and where it lands the browser."""
|
||||
|
||||
session_key: str
|
||||
session_key: str = field(repr=False)
|
||||
claims: UiSessionClaims
|
||||
redirect_url: str
|
||||
|
||||
|
|
@ -106,7 +108,10 @@ class DashboardSession:
|
|||
@dataclass(frozen=True, slots=True)
|
||||
class ManagementClient:
|
||||
proxy: ProxyClient
|
||||
master_key: str
|
||||
master_key: str = field(repr=False)
|
||||
|
||||
def with_caller(self, caller: Caller) -> ManagementClient:
|
||||
return replace(self, proxy=self.proxy.with_caller(caller))
|
||||
|
||||
def llm_only_key(self) -> str:
|
||||
return self.proxy.generate_key(KeyGenerateBody(models=[], allowed_routes=["llm_api_routes"]))
|
||||
|
|
@ -117,7 +122,7 @@ class ManagementClient:
|
|||
dashboard creates it under the session key their sign-in minted). Returns
|
||||
the outcome rather than unwrapping it, so a caller can poll a route that is
|
||||
only transiently refusing."""
|
||||
headers = self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key)
|
||||
headers = self.proxy.management_headers(caller_key)
|
||||
return self.proxy.transport.post(
|
||||
"/key/generate",
|
||||
headers=headers,
|
||||
|
|
@ -131,9 +136,9 @@ class ManagementClient:
|
|||
sign-in minted, never the master key). Returns the outcome rather than
|
||||
unwrapping it, so a caller can poll a route that is only transiently
|
||||
refusing; `update_key_models` is the unwrapping shorthand."""
|
||||
headers = self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key)
|
||||
headers = self.proxy.management_headers(caller_key)
|
||||
last: Result[NoBody] = NetworkError(message="/key/update was never attempted")
|
||||
for attempt in range(_KEY_WRITE_ATTEMPTS):
|
||||
for attempt in range(retry_attempts(_KEY_WRITE_ATTEMPTS)):
|
||||
last = self.proxy.transport.post(
|
||||
"/key/update",
|
||||
headers=headers,
|
||||
|
|
@ -144,6 +149,7 @@ class ManagementClient:
|
|||
case UnknownApiError(body=error_body) if any(
|
||||
marker in error_body.lower() for marker in _TRANSIENT_BACKEND_MARKERS
|
||||
):
|
||||
warnings.warn(f"Transient backend response on attempt {attempt + 1}", RuntimeWarning, stacklevel=2)
|
||||
time.sleep(0.5 * (attempt + 1))
|
||||
continue
|
||||
case _:
|
||||
|
|
@ -153,25 +159,26 @@ class ManagementClient:
|
|||
def update_key_models(self, key: str, models: list[str]) -> None:
|
||||
_ = unwrap(self.update_key(KeyUpdateBody(key=key, models=models)))
|
||||
|
||||
def key_info_as(self, key: str, *, caller_key: str) -> Result[KeyInfoResponse]:
|
||||
def key_info_as(self, key: str, *, caller_key: str | None = None) -> Result[KeyInfoResponse]:
|
||||
return self.proxy.transport.get(
|
||||
"/key/info",
|
||||
headers=self.proxy.transport.bearer(caller_key),
|
||||
headers=self.proxy.management_headers(caller_key),
|
||||
params=KeyInfoParams(key=key),
|
||||
response_type=KeyInfoResponse,
|
||||
)
|
||||
|
||||
def delete_key_strict(self, key: str, *, caller_key: str | None = None) -> None:
|
||||
def delete_key_strict(self, key: str, *, caller_key: str | None = None, missing_ok: bool = False) -> None:
|
||||
"""Strict delete for the act phase of a test: a failed delete is a hard
|
||||
failure, unlike the warn-only ProxyClient.delete_key used at teardown."""
|
||||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/key/delete",
|
||||
headers=self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key),
|
||||
json=KeyDeleteBody(keys=[key]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
result = self.proxy.transport.post(
|
||||
"/key/delete",
|
||||
headers=self.proxy.management_headers(caller_key),
|
||||
json=KeyDeleteBody(keys=[key]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
if missing_ok and isinstance(result, UnknownApiError) and result.status_code == 404:
|
||||
return
|
||||
_ = unwrap(result)
|
||||
|
||||
def delete_model_strict(self, model_id: str) -> None:
|
||||
"""Strict delete for the act phase of a test: a failed delete is a hard
|
||||
|
|
@ -179,7 +186,7 @@ class ManagementClient:
|
|||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/model/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=ModelDeleteBody(id=model_id),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -190,7 +197,7 @@ class ManagementClient:
|
|||
Connection button, probing the live provider with the supplied params."""
|
||||
return self.proxy.transport.post(
|
||||
"/health/test_connection",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=ConnectionTestResponse,
|
||||
timeout=120.0,
|
||||
|
|
@ -200,7 +207,7 @@ class ManagementClient:
|
|||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/key/block",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=KeyBlockBody(key=key),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -209,7 +216,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/key/regenerate",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=KeyRegenerateBody(key=key, grace_period=grace_period),
|
||||
response_type=KeyGenerateResponse,
|
||||
)
|
||||
|
|
@ -219,7 +226,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
f"/key/{key}/reset_spend",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=KeyResetSpendBody(reset_to=reset_to),
|
||||
response_type=KeyResetSpendResponse,
|
||||
)
|
||||
|
|
@ -228,7 +235,7 @@ class ManagementClient:
|
|||
def key_list(self, key_alias: str, *, caller_key: str | None = None) -> Result[KeyListResponse]:
|
||||
"""GET /key/list, the Virtual Keys page's own inventory call. `caller_key` is
|
||||
who is asking: the master key by default, or a virtual key."""
|
||||
headers = self.proxy.transport.master if caller_key is None else self.proxy.transport.bearer(caller_key)
|
||||
headers = self.proxy.management_headers(caller_key)
|
||||
return self.proxy.transport.get(
|
||||
"/key/list",
|
||||
headers=headers,
|
||||
|
|
@ -266,7 +273,7 @@ class ManagementClient:
|
|||
team_id = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/team/new",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=TeamNewResponse,
|
||||
)
|
||||
|
|
@ -276,10 +283,10 @@ class ManagementClient:
|
|||
|
||||
def update_team(self, body: TeamUpdateBody) -> None:
|
||||
last: Result[NoBody] | None = None
|
||||
for attempt in range(5):
|
||||
for attempt in range(retry_attempts(5)):
|
||||
last = self.proxy.transport.post(
|
||||
"/team/update",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -289,6 +296,7 @@ class ManagementClient:
|
|||
case UnknownApiError(body=body_text) if (
|
||||
"connecting to redis" in body_text.lower() or "name resolution" in body_text.lower()
|
||||
):
|
||||
warnings.warn(f"Transient backend response on attempt {attempt + 1}", RuntimeWarning, stacklevel=2)
|
||||
time.sleep(0.5 * (attempt + 1))
|
||||
continue
|
||||
case _:
|
||||
|
|
@ -299,7 +307,7 @@ class ManagementClient:
|
|||
def delete_team(self, team_id: str) -> None:
|
||||
_ = self.proxy.transport.post(
|
||||
"/team/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=TeamDeleteBody(team_ids=[team_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -308,7 +316,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/team/info",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=TeamInfoParams(team_id=team_id),
|
||||
response_type=TeamInfoResponse,
|
||||
)
|
||||
|
|
@ -320,7 +328,7 @@ class ManagementClient:
|
|||
for entry in unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/team/list",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=NoBody(),
|
||||
response_type=TeamListResponse,
|
||||
)
|
||||
|
|
@ -328,14 +336,16 @@ class ManagementClient:
|
|||
)
|
||||
|
||||
def team_info_status(self, team_id: str) -> ProbeResult:
|
||||
return self.proxy.transport.probe("/team/info", params=TeamInfoParams(team_id=team_id))
|
||||
return self.proxy.transport.probe(
|
||||
"/team/info", params=TeamInfoParams(team_id=team_id), headers=self.proxy.management_headers()
|
||||
)
|
||||
|
||||
def _wait_for_team(self, team_id: str) -> None:
|
||||
last: Result[TeamInfoResponse] | None = None
|
||||
for _ in range(_TEAM_READY_ATTEMPTS):
|
||||
for _ in range(retry_attempts(_TEAM_READY_ATTEMPTS)):
|
||||
last = self.proxy.transport.get(
|
||||
"/team/info",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=TeamInfoParams(team_id=team_id),
|
||||
response_type=TeamInfoResponse,
|
||||
)
|
||||
|
|
@ -343,25 +353,29 @@ class ManagementClient:
|
|||
case Success():
|
||||
return
|
||||
case _:
|
||||
warnings.warn("Repeating team read while the team becomes available", RuntimeWarning, stacklevel=2)
|
||||
time.sleep(_TEAM_READY_SLEEP_SECONDS)
|
||||
assert last is not None
|
||||
raise AssertionError(last)
|
||||
|
||||
def add_team_member(self, team_id: str, user_id: str) -> None:
|
||||
last: Result[NoBody] | None = None
|
||||
for attempt in range(_TEAM_READY_ATTEMPTS):
|
||||
for attempt in range(retry_attempts(_TEAM_READY_ATTEMPTS)):
|
||||
last = self.proxy.transport.post(
|
||||
"/team/member_add",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=TeamMemberAddBody(team_id=team_id, member=TeamMemberEntry(role="user", user_id=user_id)),
|
||||
response_type=NoBody,
|
||||
)
|
||||
match last:
|
||||
case Success():
|
||||
return
|
||||
case UnknownApiError(body=body) if (
|
||||
"doesn't exist" in body and attempt + 1 < _TEAM_READY_ATTEMPTS
|
||||
case UnknownApiError(body=body) if "doesn't exist" in body and attempt + 1 < retry_attempts(
|
||||
_TEAM_READY_ATTEMPTS
|
||||
):
|
||||
warnings.warn(
|
||||
"Retrying team membership while the team becomes available", RuntimeWarning, stacklevel=2
|
||||
)
|
||||
time.sleep(_TEAM_READY_SLEEP_SECONDS)
|
||||
continue
|
||||
case _:
|
||||
|
|
@ -373,7 +387,7 @@ class ManagementClient:
|
|||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/team/member_delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=TeamMemberDeleteBody(team_id=team_id, user_id=user_id),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -383,7 +397,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/user/new",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=UserNewResponse,
|
||||
)
|
||||
|
|
@ -393,7 +407,7 @@ class ManagementClient:
|
|||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/customer/new",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=CustomerNewBody(user_id=user_id),
|
||||
response_type=CustomerResponse,
|
||||
)
|
||||
|
|
@ -404,7 +418,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/customer/info",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=CustomerInfoParams(end_user_id=end_user_id),
|
||||
response_type=CustomerResponse,
|
||||
)
|
||||
|
|
@ -413,7 +427,7 @@ class ManagementClient:
|
|||
def delete_customer(self, user_id: str) -> None:
|
||||
_ = self.proxy.transport.post(
|
||||
"/customer/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=CustomerDeleteBody(user_ids=[user_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -422,7 +436,7 @@ class ManagementClient:
|
|||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/user/update",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -431,7 +445,7 @@ class ManagementClient:
|
|||
def delete_user(self, user_id: str) -> None:
|
||||
_ = self.proxy.transport.post(
|
||||
"/user/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=UserDeleteBody(user_ids=[user_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -442,17 +456,17 @@ class ManagementClient:
|
|||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/user/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=UserDeleteBody(user_ids=[user_id]),
|
||||
response_type=UserDeleteResponse,
|
||||
)
|
||||
)
|
||||
|
||||
def user_info(self, user_id: str) -> UserInfoResponse:
|
||||
def user_info(self, user_id: str | None = None) -> UserInfoResponse:
|
||||
return unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/user/info",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=UserInfoParams(user_id=user_id),
|
||||
response_type=UserInfoResponse,
|
||||
)
|
||||
|
|
@ -462,7 +476,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/user/list",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=UserListParams(user_ids=user_id),
|
||||
response_type=UserListResponse,
|
||||
)
|
||||
|
|
@ -472,7 +486,7 @@ class ManagementClient:
|
|||
listing = unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/user/list",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=UserListParams(user_ids=user_id),
|
||||
response_type=UserListResponse,
|
||||
)
|
||||
|
|
@ -483,7 +497,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/organization/new",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=OrgNewResponse,
|
||||
)
|
||||
|
|
@ -493,7 +507,7 @@ class ManagementClient:
|
|||
_ = unwrap(
|
||||
self.proxy.transport.patch(
|
||||
"/organization/update",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -502,7 +516,7 @@ class ManagementClient:
|
|||
def delete_org(self, organization_id: str) -> None:
|
||||
_ = self.proxy.transport.delete(
|
||||
"/organization/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=OrgDeleteBody(organization_ids=[organization_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -511,19 +525,24 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/organization/info",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=OrgInfoParams(organization_id=organization_id),
|
||||
response_type=OrgInfoResponse,
|
||||
)
|
||||
)
|
||||
|
||||
def org_info_status(self, organization_id: str) -> ProbeResult:
|
||||
return self.proxy.transport.probe("/organization/info", params=OrgInfoParams(organization_id=organization_id))
|
||||
return self.proxy.transport.probe(
|
||||
"/organization/info",
|
||||
params=OrgInfoParams(organization_id=organization_id),
|
||||
headers=self.proxy.management_headers(),
|
||||
)
|
||||
|
||||
def create_tag(self, body: TagNewBody) -> None:
|
||||
_ = unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/tag/new",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -532,7 +551,7 @@ class ManagementClient:
|
|||
def delete_tag(self, name: str) -> None:
|
||||
_ = self.proxy.transport.post(
|
||||
"/tag/delete",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=TagDeleteBody(name=name),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -542,7 +561,7 @@ class ManagementClient:
|
|||
unwrap(
|
||||
self.proxy.transport.get(
|
||||
"/tag/list",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
params=NoBody(),
|
||||
response_type=TagListResponse,
|
||||
)
|
||||
|
|
@ -553,7 +572,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.post(
|
||||
"/v1/mcp/server",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=McpServerRow,
|
||||
)
|
||||
|
|
@ -565,7 +584,7 @@ class ManagementClient:
|
|||
return unwrap(
|
||||
self.proxy.transport.put(
|
||||
"/v1/mcp/server",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=body,
|
||||
response_type=McpServerRow,
|
||||
)
|
||||
|
|
@ -576,7 +595,7 @@ class ManagementClient:
|
|||
unwrap it while a deferred teardown can ignore an already-deleted server."""
|
||||
return self.proxy.transport.delete(
|
||||
f"/v1/mcp/server/{server_id}",
|
||||
headers=self.proxy.transport.master,
|
||||
headers=self.proxy.management_headers(),
|
||||
json=NoBody(),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2,60 +2,247 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from e2e_config import CHEAP_OPENAI_MODEL, unique_marker
|
||||
from e2e_config import CHEAP_OPENAI_MODEL, PROXY_BASE_URL, unique_marker
|
||||
from e2e_http import UnauthorizedError, UnknownApiError, unwrap
|
||||
from idp import ADMIN_CLIENT_ID, Identity, Keycloak
|
||||
from idp import ADMIN_CLIENT_ID, Identity, Keycloak, token_claims
|
||||
from lifecycle import ResourceManager
|
||||
from management.jwt_actors import ActorFactory, ActorRole
|
||||
from management_client import ManagementClient
|
||||
from models import KeyGenerateBody, KeyUpdateBody, TeamNewBody, UserNewBody
|
||||
from models import KeyGenerateBody, KeyUpdateBody, TeamNewBody, UserInfoParams, UserInfoResponse, UserNewBody
|
||||
from proxy_client import Caller
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
class TestJwtManagement:
|
||||
@pytest.mark.covers("mgmt.key.jwt.lifecycle")
|
||||
def test_admin_creates_reads_updates_clears_and_deletes_a_key(
|
||||
self, client: ManagementClient, idp: Keycloak, jwt_identity: Identity, resources: ResourceManager
|
||||
) -> None:
|
||||
admin: Final = idp.access_token(jwt_identity, client_id=ADMIN_CLIENT_ID)
|
||||
alias: Final = f"e2e-jwt-key-{unique_marker()}"
|
||||
created: Final = unwrap(
|
||||
client.generate_key(
|
||||
KeyGenerateBody(key_alias=alias, team_id=jwt_identity.group, models=[CHEAP_OPENAI_MODEL]),
|
||||
caller_key=admin,
|
||||
@pytest.mark.parametrize(
|
||||
"role",
|
||||
(
|
||||
"proxy_admin",
|
||||
"proxy_admin_viewer",
|
||||
"organization_admin",
|
||||
"team_admin",
|
||||
"team_member",
|
||||
"internal_user",
|
||||
"internal_user_viewer",
|
||||
"unrelated_user",
|
||||
),
|
||||
)
|
||||
@pytest.mark.covers("mgmt.user.jwt.database_roles")
|
||||
def test_actor_subject_and_database_role(self, actor_factory: ActorFactory, role: ActorRole) -> None:
|
||||
tenants: Final = (
|
||||
(actor_factory.tenant(),) if role in ("organization_admin", "team_admin", "team_member") else ()
|
||||
)
|
||||
actor: Final = actor_factory.create(role, tenants=tenants)
|
||||
caller: Final = actor.mint_caller(actor_factory.idp)
|
||||
claims: Final = token_claims(caller.credential)
|
||||
assert claims.sub == actor.identity.user_id
|
||||
assert claims.iss == actor_factory.idp.issuer
|
||||
assert claims.aud == "litellm-e2e" or "litellm-e2e" in claims.aud
|
||||
assert actor.identity.groups == ()
|
||||
assert ("litellm_proxy_admin" in claims.scope.split()) == (role == "proxy_admin")
|
||||
stored: Final = actor_factory.bootstrap.user_info(actor.identity.user_id)
|
||||
assert stored.user_id == actor.identity.user_id
|
||||
assert stored.user_info.user_role == actor.global_role
|
||||
bound: Final = actor_factory.bootstrap.with_caller(caller)
|
||||
own: Final = unwrap(
|
||||
bound.proxy.transport.get(
|
||||
"/user/info",
|
||||
headers=bound.proxy.management_headers(),
|
||||
params=UserInfoParams(),
|
||||
response_type=UserInfoResponse,
|
||||
)
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(created.key))
|
||||
assert own.user_id == actor.identity.user_id
|
||||
assert own.user_info.user_role == actor.global_role
|
||||
for tenant in tenants:
|
||||
info = actor_factory.bootstrap.team_info(tenant.team_id)
|
||||
assert info.organization_id == tenant.organization_id
|
||||
assert {(member.user_id, member.role) for member in info.members_with_roles} == {
|
||||
(actor.identity.user_id, "admin" if role == "team_admin" else "user")
|
||||
}
|
||||
assert {
|
||||
(member.user_id, member.user_role)
|
||||
for member in actor_factory.bootstrap.org_info(tenant.organization_id).members
|
||||
} == {(actor.identity.user_id, "org_admin" if role == "organization_admin" else "internal_user")}
|
||||
|
||||
original: Final = unwrap(client.key_info_as(created.key, caller_key=admin)).info
|
||||
assert original.key_alias == alias and original.team_id == jwt_identity.group
|
||||
@pytest.mark.covers("mgmt.key.jwt.viewer_denied")
|
||||
def test_admin_viewer_reads_but_cannot_update(self, actor_factory: ActorFactory) -> None:
|
||||
actor: Final = actor_factory.create("proxy_admin_viewer")
|
||||
viewer: Final = actor_factory.bootstrap.with_caller(actor.mint_caller(actor_factory.idp))
|
||||
alias: Final = f"e2e-viewer-{unique_marker()}"
|
||||
key: Final = actor_factory.key().key
|
||||
unwrap(actor_factory.bootstrap.update_key(KeyUpdateBody(key=key, key_alias=alias)))
|
||||
assert viewer.proxy.key_info(key).key_alias == alias
|
||||
denied: Final = viewer.update_key(KeyUpdateBody(key=key, key_alias="forbidden"))
|
||||
assert isinstance(denied, UnknownApiError) and denied.status_code == 403, f"viewer write was accepted: {denied}"
|
||||
assert "proxy_admin_viewer" in denied.body and "/key/update" in denied.body
|
||||
assert actor_factory.bootstrap.proxy.key_info(key).key_alias == alias
|
||||
|
||||
@pytest.mark.covers("mgmt.user.oidc.identity_mapping")
|
||||
def test_oidc_browser_profile_identity_mapping(self, actor_factory: ActorFactory) -> None:
|
||||
actor: Final = actor_factory.create("internal_user")
|
||||
idp: Final = actor_factory.idp.with_strict_cleanup()
|
||||
discovery: Final = idp.discovery()
|
||||
assert discovery.issuer == idp.issuer
|
||||
assert discovery.jwks_uri == idp.jwks_url
|
||||
callback: Final = f"{PROXY_BASE_URL}/sso/callback"
|
||||
browser: Final = idp.browser_client(callback_url=callback, defer=actor_factory.resources.defer)
|
||||
token: Final = idp.browser_token(actor.identity, browser)
|
||||
assert token_claims(token).sub == actor.identity.user_id
|
||||
userinfo: Final = idp.userinfo(token)
|
||||
assert userinfo.sub == actor.identity.user_id
|
||||
assert userinfo.email == f"{actor.identity.username}@example.com"
|
||||
assert browser.environment(discovery)["GENERIC_USER_ID_ATTRIBUTE"] == "sub"
|
||||
|
||||
@pytest.mark.covers("mgmt.key.jwt.lifecycle")
|
||||
@pytest.mark.parametrize("credential_kind", ("direct_jwt", "virtual_key"))
|
||||
def test_admin_creates_reads_updates_clears_and_deletes_a_key(
|
||||
self,
|
||||
actor_factory: ActorFactory,
|
||||
credential_kind: Literal["direct_jwt", "virtual_key"],
|
||||
) -> None:
|
||||
tenant: Final = actor_factory.tenant()
|
||||
actor: Final = actor_factory.create("proxy_admin", tenants=(tenant,), profile="group_scoped")
|
||||
virtual_key: Final = (
|
||||
actor_factory.key(user_id=actor.identity.user_id).key if credential_kind == "virtual_key" else None
|
||||
)
|
||||
admin: Final = virtual_key if virtual_key is not None else actor.mint_caller(actor_factory.idp).credential
|
||||
bound: Final = actor_factory.bootstrap.with_caller(
|
||||
Caller(credential=admin, kind=credential_kind, role="proxy_admin")
|
||||
)
|
||||
assert bound.user_info().user_id == actor.identity.user_id
|
||||
alias: Final = f"e2e-jwt-key-{unique_marker()}"
|
||||
created: Final = unwrap(
|
||||
bound.generate_key(
|
||||
KeyGenerateBody(key_alias=alias, team_id=tenant.team_id, models=[CHEAP_OPENAI_MODEL]),
|
||||
)
|
||||
)
|
||||
actor_factory.resources.defer(lambda: actor_factory.bootstrap.delete_key_strict(created.key, missing_ok=True))
|
||||
|
||||
original: Final = unwrap(bound.key_info_as(created.key)).info
|
||||
assert original.key_alias == alias and original.team_id == tenant.team_id
|
||||
assert original.models == [CHEAP_OPENAI_MODEL]
|
||||
|
||||
updated_alias: Final = f"{alias}-updated"
|
||||
unwrap(
|
||||
client.update_key(KeyUpdateBody(key=created.key, key_alias=updated_alias, rpm_limit=120), caller_key=admin)
|
||||
)
|
||||
updated: Final = unwrap(client.key_info_as(created.key, caller_key=admin)).info
|
||||
unwrap(bound.update_key(KeyUpdateBody(key=created.key, key_alias=updated_alias, rpm_limit=120)))
|
||||
updated: Final = unwrap(bound.key_info_as(created.key)).info
|
||||
assert updated.key_alias == updated_alias and updated.rpm_limit == 120
|
||||
assert updated.models == [CHEAP_OPENAI_MODEL], "omitted models must preserve the restriction"
|
||||
|
||||
unwrap(client.update_key(KeyUpdateBody(key=created.key, models=[]), caller_key=admin))
|
||||
cleared: Final = unwrap(client.key_info_as(created.key, caller_key=admin)).info
|
||||
unwrap(bound.update_key(KeyUpdateBody(key=created.key, models=[])))
|
||||
cleared: Final = unwrap(bound.key_info_as(created.key)).info
|
||||
assert cleared.models == [] and cleared.rpm_limit == 120
|
||||
|
||||
assert unwrap(client.key_list(updated_alias, caller_key=admin)).total_count == 1
|
||||
client.delete_key_strict(created.key, caller_key=admin)
|
||||
assert unwrap(client.key_list(updated_alias, caller_key=admin)).total_count == 0
|
||||
assert unwrap(bound.key_list(updated_alias)).total_count == 1
|
||||
bound.delete_key_strict(created.key)
|
||||
assert unwrap(bound.key_list(updated_alias)).total_count == 0
|
||||
|
||||
@pytest.mark.covers("mgmt.team.jwt.tenant_isolation")
|
||||
def test_two_actor_sets_keep_tenants_and_keys_isolated(self, actor_factory: ActorFactory) -> None:
|
||||
first: Final = actor_factory.tenant()
|
||||
second: Final = actor_factory.tenant()
|
||||
assert first.organization_id != second.organization_id and first.team_id != second.team_id
|
||||
actors: Final = tuple(
|
||||
actor_factory.create("team_member", tenants=(tenant,), profile="group_scoped") for tenant in (first, second)
|
||||
)
|
||||
assert actors[0].identity.user_id != actors[1].identity.user_id
|
||||
callers: Final = tuple(
|
||||
actor_factory.bootstrap.with_caller(actor.mint_caller(actor_factory.idp)) for actor in actors
|
||||
)
|
||||
keys: Final = tuple(actor_factory.key(tenant) for tenant in (first, second))
|
||||
assert keys[0].key != keys[1].key
|
||||
assert callers[0].proxy.key_info(keys[0].key).team_id == first.team_id
|
||||
assert callers[1].proxy.key_info(keys[1].key).team_id == second.team_id
|
||||
for caller, other_key in ((callers[0], keys[1].key), (callers[1], keys[0].key)):
|
||||
hidden = caller.key_info_as(other_key)
|
||||
assert isinstance(hidden, UnknownApiError) and hidden.status_code == 403
|
||||
assert tuple(actor.identity.groups for actor in actors) == ((first.team_id,), (second.team_id,))
|
||||
|
||||
@pytest.mark.covers("mgmt.team.jwt.multiple_memberships")
|
||||
def test_multi_group_actor_keeps_exact_memberships(self, actor_factory: ActorFactory) -> None:
|
||||
tenants: Final = (actor_factory.tenant(), actor_factory.tenant())
|
||||
actor: Final = actor_factory.create("team_member", tenants=tenants, profile="group_scoped")
|
||||
claims: Final = token_claims(actor.mint_caller(actor_factory.idp).credential)
|
||||
assert set(claims.groups) == {tenant.team_id for tenant in tenants}
|
||||
assert "litellm_proxy_admin" not in claims.scope.split()
|
||||
assert actor.identity.groups == tuple(tenant.team_id for tenant in tenants)
|
||||
for tenant in tenants:
|
||||
assert {
|
||||
(entry.user_id, entry.role)
|
||||
for entry in actor_factory.bootstrap.team_info(tenant.team_id).members_with_roles
|
||||
} == {(actor.identity.user_id, "user")}
|
||||
|
||||
@pytest.mark.covers("mgmt.user.jwt.cleanup")
|
||||
def test_successful_actor_cleanup_removes_owned_state(self, actor_factory: ActorFactory) -> None:
|
||||
resources: Final = ResourceManager(client=actor_factory.bootstrap.proxy, strict_cleanup=True)
|
||||
factory: Final = ActorFactory(bootstrap=actor_factory.bootstrap, idp=actor_factory.idp, resources=resources)
|
||||
try:
|
||||
tenant: Final = factory.tenant()
|
||||
actor: Final = factory.create("team_member", tenants=(tenant,), profile="group_scoped")
|
||||
key: Final = factory.key(tenant)
|
||||
alias: Final = factory.bootstrap.proxy.key_info(key.key).key_alias
|
||||
assert alias is not None
|
||||
finally:
|
||||
resources.teardown()
|
||||
assert factory.bootstrap.user_count(actor.identity.user_id) == 0
|
||||
assert factory.bootstrap.key_alias_count(alias) == 0
|
||||
assert factory.bootstrap.team_info_status(tenant.team_id).status_code == 404
|
||||
assert factory.bootstrap.org_info_status(tenant.organization_id).status_code == 404
|
||||
factory.idp.assert_absent("users", actor.identity.user_id)
|
||||
factory.idp.assert_absent("groups", tenant.group_id)
|
||||
|
||||
@pytest.mark.parametrize("stage", ("group", "user"))
|
||||
@pytest.mark.covers("mgmt.user.jwt.partial_cleanup")
|
||||
def test_partial_setup_removes_previously_created_identities(
|
||||
self,
|
||||
actor_factory: ActorFactory,
|
||||
stage: Literal["group", "user"],
|
||||
) -> None:
|
||||
idp: Final = actor_factory.idp.with_strict_cleanup()
|
||||
resources: Final = ResourceManager(client=actor_factory.bootstrap.proxy, strict_cleanup=True)
|
||||
marker: Final = unique_marker()
|
||||
group_id: Final = idp.create_group(f"e2e-partial-{marker}")
|
||||
resources.defer(lambda: idp.delete_group(group_id))
|
||||
try:
|
||||
identity: Final = (
|
||||
idp.provision_user(
|
||||
marker=marker,
|
||||
groups=(f"e2e-partial-{marker}",),
|
||||
group_ids=(group_id,),
|
||||
defer=resources.defer,
|
||||
)
|
||||
if stage == "user"
|
||||
else None
|
||||
)
|
||||
if identity is None:
|
||||
with pytest.raises(pytest.fail.Exception, match="HTTP 409"):
|
||||
idp.create_group(f"e2e-partial-{marker}")
|
||||
else:
|
||||
with pytest.raises(pytest.fail.Exception, match="HTTP 409"):
|
||||
idp.create_user(
|
||||
username=identity.username,
|
||||
email=f"{identity.username}@example.com",
|
||||
password=identity.password,
|
||||
groups=identity.groups,
|
||||
)
|
||||
finally:
|
||||
resources.teardown()
|
||||
idp.assert_absent("groups", group_id)
|
||||
if identity is not None:
|
||||
idp.assert_absent("users", identity.user_id)
|
||||
|
||||
@pytest.mark.covers("mgmt.key.jwt.member_denied", "mgmt.key.jwt.other_team_denied")
|
||||
def test_member_cannot_write_and_another_team_cannot_read_the_key(
|
||||
self, client: ManagementClient, idp: Keycloak, jwt_identity: Identity, resources: ResourceManager
|
||||
) -> None:
|
||||
admin: Final = idp.access_token(jwt_identity, client_id=ADMIN_CLIENT_ID)
|
||||
bound: Final = client.with_caller(Caller(credential=admin, kind="direct_jwt", role="proxy_admin"))
|
||||
member: Final = idp.access_token(jwt_identity)
|
||||
member_client: Final = client.with_caller(Caller(credential=member, kind="direct_jwt", role="team_member"))
|
||||
alias: Final = f"e2e-jwt-owned-{unique_marker()}"
|
||||
created: Final = unwrap(
|
||||
client.generate_key(KeyGenerateBody(key_alias=alias, team_id=jwt_identity.group), caller_key=admin)
|
||||
|
|
@ -63,14 +250,14 @@ class TestJwtManagement:
|
|||
resources.defer(lambda: client.proxy.delete_key(created.key))
|
||||
|
||||
client.add_team_member(jwt_identity.group, jwt_identity.user_id)
|
||||
assert unwrap(client.key_info_as(created.key, caller_key=member)).info.key_alias == alias
|
||||
assert unwrap(member_client.key_info_as(created.key)).info.key_alias == alias
|
||||
|
||||
refused: Final = client.update_key(KeyUpdateBody(key=created.key, key_alias="forbidden"), caller_key=member)
|
||||
refused: Final = member_client.update_key(KeyUpdateBody(key=created.key, key_alias="forbidden"))
|
||||
assert isinstance(refused, UnauthorizedError), f"member write was accepted: {refused}"
|
||||
assert "does not have permissions for endpoint" in refused.body.lower(), (
|
||||
f"expected a permission denial: {refused}"
|
||||
)
|
||||
assert unwrap(client.key_info_as(created.key, caller_key=admin)).info.key_alias == alias
|
||||
assert unwrap(bound.key_info_as(created.key)).info.key_alias == alias
|
||||
|
||||
marker: Final = unique_marker()
|
||||
outsider: Final = idp.provision(marker=marker, group=f"e2e-jwt-team-{marker}", defer=resources.defer)
|
||||
|
|
@ -88,4 +275,4 @@ class TestJwtManagement:
|
|||
assert isinstance(hidden, UnknownApiError) and hidden.status_code == 403, (
|
||||
f"another team must not read this key: {hidden}"
|
||||
)
|
||||
assert unwrap(client.key_info_as(created.key, caller_key=admin)).info.team_id == jwt_identity.group
|
||||
assert unwrap(bound.key_info_as(created.key)).info.team_id == jwt_identity.group
|
||||
|
|
|
|||
|
|
@ -1091,13 +1091,13 @@ class UiLoginBody(BaseModel):
|
|||
|
||||
|
||||
class UiLoginResponse(BaseModel):
|
||||
token: str
|
||||
token: str = Field(repr=False)
|
||||
redirect_url: str
|
||||
|
||||
|
||||
class UiSessionClaims(BaseModel):
|
||||
user_id: str
|
||||
key: str
|
||||
key: str = Field(repr=False)
|
||||
user_role: str
|
||||
login_method: Literal["sso", "username_password"]
|
||||
exp: int
|
||||
|
|
@ -1135,6 +1135,7 @@ class TeamInfoParams(BaseModel):
|
|||
|
||||
|
||||
class TeamData(BaseModel):
|
||||
organization_id: str | None = None
|
||||
team_alias: str | None = None
|
||||
models: list[str] = []
|
||||
members_with_roles: list[TeamMemberEntry] = []
|
||||
|
|
@ -1175,6 +1176,7 @@ class UserNewBody(BaseModel):
|
|||
user_email: str
|
||||
user_role: UserRole
|
||||
user_id: str | None = None
|
||||
auto_create_key: bool | None = None
|
||||
|
||||
|
||||
class UserNewResponse(BaseModel):
|
||||
|
|
@ -1187,7 +1189,7 @@ class UserUpdateBody(BaseModel):
|
|||
|
||||
|
||||
class UserInfoParams(BaseModel):
|
||||
user_id: str
|
||||
user_id: str | None = None
|
||||
|
||||
|
||||
class UserData(BaseModel):
|
||||
|
|
@ -1240,16 +1242,36 @@ class OrgInfoParams(BaseModel):
|
|||
organization_id: str
|
||||
|
||||
|
||||
class OrgMembership(BaseModel):
|
||||
user_id: str
|
||||
user_role: str
|
||||
|
||||
|
||||
class OrgInfoResponse(BaseModel):
|
||||
organization_id: str
|
||||
organization_alias: str | None = None
|
||||
models: list[str] = []
|
||||
members: tuple[OrgMembership, ...] = ()
|
||||
|
||||
|
||||
class OrgMemberEntry(BaseModel):
|
||||
user_id: str
|
||||
role: Literal["org_admin", "internal_user"]
|
||||
|
||||
|
||||
class OrgMemberAddBody(BaseModel):
|
||||
organization_id: str
|
||||
member: OrgMemberEntry
|
||||
|
||||
|
||||
class OrgDeleteBody(BaseModel):
|
||||
organization_ids: list[str]
|
||||
|
||||
|
||||
class OrgDeleteResponse(RootModel[tuple[OrgInfoResponse, ...]]):
|
||||
pass
|
||||
|
||||
|
||||
# ---------- tags (management) ----------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -92,6 +92,7 @@ from fixture_mode import (
|
|||
current_test_key,
|
||||
parse_fixture_mode,
|
||||
)
|
||||
from fixture_profile import IneligibleRequest, MatchProfile, match_profile, strict_identity
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
EDGE_MOUNTS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
|
|
@ -404,6 +405,12 @@ def _miss_message(test_key: str, slug: str, canonical: CanonicalRequest, bundle:
|
|||
f"under {slug}; re-record with E2E_FIXTURE_MODE=record"
|
||||
)
|
||||
closest, closest_file = _closest_recorded(canonical, recorded)
|
||||
if bundle.manifest.match_profile == "stateless_v1":
|
||||
expected: Final = _JSON.validate_json(closest.content)
|
||||
actual: Final = _JSON.validate_json(canonical.content)
|
||||
assert isinstance(expected, dict) and isinstance(actual, dict)
|
||||
changed: Final = ", ".join(key for key in expected if expected[key] != actual.get(key))
|
||||
return f"stateless_v1 replay mismatch: {changed or 'method/path'}; re-record with E2E_FIXTURE_MODE=record"
|
||||
diff: Final = "\n".join(
|
||||
islice(
|
||||
difflib.unified_diff(
|
||||
|
|
@ -785,11 +792,33 @@ def handle_edge_request(
|
|||
mount, _, upstream_path = split.path.lstrip("/").partition("/")
|
||||
upstream_base: Final = mounts.get(mount)
|
||||
if upstream_base is None:
|
||||
return _text_reply(
|
||||
404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}"
|
||||
return _text_reply(404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}")
|
||||
profile: Final = (
|
||||
backend.recorder.profile
|
||||
if isinstance(backend, RecordEdge)
|
||||
else backend.source.bundle.manifest.match_profile
|
||||
if isinstance(backend, ReplayEdge)
|
||||
else "legacy"
|
||||
)
|
||||
identity: Final = (
|
||||
strict_identity(
|
||||
method=method,
|
||||
path=split.path,
|
||||
query=split.query,
|
||||
headers=headers,
|
||||
body=body,
|
||||
mount=mount,
|
||||
upstream_base=upstream_base,
|
||||
)
|
||||
request: Final = edge_request(
|
||||
method, split.path, split.query, body, _header_value(headers, "content-type")
|
||||
if profile == "stateless_v1"
|
||||
else None
|
||||
)
|
||||
if isinstance(identity, IneligibleRequest):
|
||||
return _text_reply(REPLAY_MISS_STATUS, f"stateless_v1 eligibility error: {identity.reason}")
|
||||
request: Final = (
|
||||
RecordedRequest(method=method.lower(), path=split.path, headers={}, strict_identity=identity)
|
||||
if identity is not None
|
||||
else edge_request(method, split.path, split.query, body, _header_value(headers, "content-type"))
|
||||
)
|
||||
match backend:
|
||||
case LiveEdge():
|
||||
|
|
@ -837,6 +866,14 @@ class _EdgeHandler(BaseHTTPRequestHandler):
|
|||
body: Final = self.rfile.read(length) if length else None
|
||||
if edge_server.observation is not None:
|
||||
edge_server.observation.observe(body)
|
||||
strict: Final = (
|
||||
isinstance(edge_server.backend, RecordEdge) and edge_server.backend.recorder.profile == "stateless_v1"
|
||||
or isinstance(edge_server.backend, ReplayEdge)
|
||||
and edge_server.backend.source.bundle.manifest.match_profile == "stateless_v1"
|
||||
)
|
||||
if strict and len({name.lower() for name in self.headers.keys()}) != len(self.headers):
|
||||
self._write_reply(_text_reply(REPLAY_MISS_STATUS, "stateless_v1 eligibility error: duplicate headers"))
|
||||
return
|
||||
outcome: Final = handle_edge_request(
|
||||
edge_server.backend,
|
||||
edge_server.mounts,
|
||||
|
|
@ -955,16 +992,16 @@ def start_provider_edge(
|
|||
|
||||
|
||||
@functools.lru_cache(maxsize=8)
|
||||
def _shared_recorder(root: Path) -> BundleRecorder:
|
||||
prepared = prepare_bundle(root)
|
||||
def _shared_recorder(root: Path, profile: MatchProfile = "legacy") -> BundleRecorder:
|
||||
prepared = prepare_bundle(root, profile=profile)
|
||||
if isinstance(prepared, UnsafeBundleDir):
|
||||
raise ValueError(f"E2E_FIXTURE_DIR {prepared.path} {prepared.reason}")
|
||||
return prepared
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=8)
|
||||
def _shared_replay_source(root: Path) -> ReplaySource:
|
||||
loaded = load_bundle(root)
|
||||
def _shared_replay_source(root: Path, profile: MatchProfile = "legacy") -> ReplaySource:
|
||||
loaded = load_bundle(root, profile=profile)
|
||||
if isinstance(loaded, UnreadableBundle):
|
||||
raise ValueError(f"cannot replay from {root}: {loaded.reason}")
|
||||
return ReplaySource(bundle=loaded)
|
||||
|
|
@ -977,11 +1014,12 @@ def _shared_edge(
|
|||
bind_host: str,
|
||||
advertise_host: str,
|
||||
forward_timeout: float,
|
||||
profile: MatchProfile,
|
||||
) -> ProviderEdge:
|
||||
backend: Final[EdgeBackend] = (
|
||||
RecordEdge(recorder=_shared_recorder(bundle_dir), lock=threading.Lock())
|
||||
RecordEdge(recorder=_shared_recorder(bundle_dir, profile), lock=threading.Lock())
|
||||
if mode == "record"
|
||||
else ReplayEdge(source=_shared_replay_source(bundle_dir))
|
||||
else ReplayEdge(source=_shared_replay_source(bundle_dir, profile))
|
||||
)
|
||||
return start_provider_edge(
|
||||
backend,
|
||||
|
|
@ -998,7 +1036,7 @@ def replay_leftover_error(*, mode_raw: str, bundle_dir: Path, test_key: str) ->
|
|||
recording it no longer matches. Inert in every other mode."""
|
||||
if parse_fixture_mode(mode_raw) != "replay":
|
||||
return None
|
||||
return _shared_replay_source(bundle_dir).leftover_error(test_key)
|
||||
return _shared_replay_source(bundle_dir, match_profile()).leftover_error(test_key)
|
||||
|
||||
|
||||
def provider_edge_api_base(
|
||||
|
|
@ -1021,10 +1059,10 @@ def provider_edge_api_base(
|
|||
return None
|
||||
case "record" | "replay":
|
||||
if mount not in EDGE_MOUNTS:
|
||||
raise ValueError(
|
||||
f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(EDGE_MOUNTS))}"
|
||||
)
|
||||
return _shared_edge(mode, bundle_dir, bind_host, advertise_host, forward_timeout).api_base(mount)
|
||||
raise ValueError(f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(EDGE_MOUNTS))}")
|
||||
return _shared_edge(mode, bundle_dir, bind_host, advertise_host, forward_timeout, match_profile()).api_base(
|
||||
mount
|
||||
)
|
||||
case _:
|
||||
assert_never(mode)
|
||||
|
||||
|
|
@ -1037,9 +1075,9 @@ def _observed_backend(mode_raw: str, bundle_dir: Path) -> EdgeBackend:
|
|||
case "live":
|
||||
return LiveEdge()
|
||||
case "record":
|
||||
return RecordEdge(_shared_recorder(bundle_dir), threading.Lock())
|
||||
return RecordEdge(_shared_recorder(bundle_dir, match_profile()), threading.Lock())
|
||||
case "replay":
|
||||
return ReplayEdge(_shared_replay_source(bundle_dir))
|
||||
return ReplayEdge(_shared_replay_source(bundle_dir, match_profile()))
|
||||
case _:
|
||||
assert_never(mode)
|
||||
|
||||
|
|
|
|||
|
|
@ -11,14 +11,23 @@ from __future__ import annotations
|
|||
import time
|
||||
import warnings
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from functools import reduce
|
||||
from dataclasses import dataclass, field, replace
|
||||
from datetime import datetime
|
||||
from functools import reduce
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
from typing import Final, Literal
|
||||
|
||||
from e2e_config import (
|
||||
CONTROL_PLANE_BASE_URL,
|
||||
MASTER_KEY,
|
||||
POLL_INTERVAL,
|
||||
POLL_TIMEOUT,
|
||||
PROXY_BASE_URL,
|
||||
PROXY_REPLICA_URLS,
|
||||
REQUEST_TIMEOUT,
|
||||
SLOW_PROVIDER_TIMEOUT_SECONDS,
|
||||
settle_propagation,
|
||||
)
|
||||
from e2e_http import (
|
||||
AnthropicHeaders,
|
||||
AuthHeaders,
|
||||
|
|
@ -55,6 +64,7 @@ from models import (
|
|||
KeyInfoParams,
|
||||
KeyInfoResponse,
|
||||
LiteLLMParamsBody,
|
||||
MemorySummaryResponse,
|
||||
ModelDeleteBody,
|
||||
ModelInfoBody,
|
||||
ModelInfoEntry,
|
||||
|
|
@ -63,7 +73,6 @@ from models import (
|
|||
ModelNewBody,
|
||||
ModelNewResponse,
|
||||
ModelsListParams,
|
||||
MemorySummaryResponse,
|
||||
ModelsListResponse,
|
||||
ModelUpdateBody,
|
||||
OcrBody,
|
||||
|
|
@ -76,23 +85,13 @@ from models import (
|
|||
TeamDeleteBody,
|
||||
TeamNewBody,
|
||||
TeamNewResponse,
|
||||
UserDeleteBody,
|
||||
UserDeleteResponse,
|
||||
ToolsetCreateBody,
|
||||
ToolsetRow,
|
||||
ToolsetUpdateBody,
|
||||
UserDeleteBody,
|
||||
UserDeleteResponse,
|
||||
)
|
||||
from e2e_config import (
|
||||
CONTROL_PLANE_BASE_URL,
|
||||
MASTER_KEY,
|
||||
POLL_INTERVAL,
|
||||
POLL_TIMEOUT,
|
||||
PROXY_BASE_URL,
|
||||
PROXY_REPLICA_URLS,
|
||||
REQUEST_TIMEOUT,
|
||||
SLOW_PROVIDER_TIMEOUT_SECONDS,
|
||||
settle_propagation,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
from transport import HttpTransport, SplitTransport, Transport, is_control_plane_path
|
||||
|
||||
RowsPredicate = Callable[[list[SpendLogRow]], bool]
|
||||
|
|
@ -421,11 +420,23 @@ def converge_timeout_message(*, what: str, replica: str, timeout: float, last_re
|
|||
)
|
||||
|
||||
|
||||
CredentialKind = Literal["master", "direct_jwt", "virtual_key", "dashboard_session"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Caller:
|
||||
credential: str = field(repr=False)
|
||||
kind: CredentialKind
|
||||
role: str
|
||||
tenant: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProxyClient:
|
||||
transport: Transport
|
||||
replicas: Mapping[str, Transport]
|
||||
control_replicas: Mapping[str, Transport]
|
||||
caller: Caller | None = None
|
||||
poll_timeout: float = 120.0
|
||||
poll_interval: float = 5.0
|
||||
model_servable_timeout: float = MODEL_SERVABLE_TIMEOUT
|
||||
|
|
@ -433,13 +444,24 @@ class ProxyClient:
|
|||
model_servable_interval: float = MODEL_SERVABLE_INTERVAL
|
||||
model_servable_request_timeout: float = MODEL_SERVABLE_REQUEST_TIMEOUT
|
||||
|
||||
def with_caller(self, caller: Caller) -> ProxyClient:
|
||||
return replace(self, caller=caller)
|
||||
|
||||
def management_headers(self, caller_key: str | None = None, *, transport: Transport | None = None) -> AuthHeaders:
|
||||
selected: Final = self.transport if transport is None else transport
|
||||
if caller_key is not None:
|
||||
return selected.bearer(caller_key)
|
||||
if self.caller is not None:
|
||||
return selected.bearer(self.caller.credential)
|
||||
return selected.master
|
||||
|
||||
# ---- keys / customers (satisfies lifecycle.ResourceClient) ----------
|
||||
|
||||
def generate_key(self, body: KeyGenerateBody) -> str:
|
||||
return unwrap(
|
||||
self.transport.post(
|
||||
"/key/generate",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=body,
|
||||
response_type=KeyGenerateResponse,
|
||||
)
|
||||
|
|
@ -448,7 +470,7 @@ class ProxyClient:
|
|||
def delete_key(self, key: str) -> None:
|
||||
_ = self.transport.post(
|
||||
"/key/delete",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=KeyDeleteBody(keys=[key]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -458,7 +480,7 @@ class ProxyClient:
|
|||
return
|
||||
_ = self.transport.post(
|
||||
"/customer/delete",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=CustomerDeleteBody(user_ids=user_ids),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -467,7 +489,7 @@ class ProxyClient:
|
|||
return unwrap(
|
||||
self.transport.get(
|
||||
"/key/info",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
params=KeyInfoParams(key=key),
|
||||
response_type=KeyInfoResponse,
|
||||
)
|
||||
|
|
@ -477,7 +499,7 @@ class ProxyClient:
|
|||
return {
|
||||
url: transport.get(
|
||||
"/debug/memory/summary",
|
||||
headers=transport.master,
|
||||
headers=self.management_headers(transport=transport),
|
||||
params=NoBody(),
|
||||
response_type=MemorySummaryResponse,
|
||||
)
|
||||
|
|
@ -524,11 +546,12 @@ class ProxyClient:
|
|||
{replica: outcome.result for replica, outcome in outcomes.items() if isinstance(outcome, Converged)}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _body_poller[R: BaseModel](
|
||||
transport: Transport, path: str, params: BaseModel, response_type: type[R]
|
||||
self, transport: Transport, path: str, params: BaseModel, response_type: type[R]
|
||||
) -> Poller[Result[R]]:
|
||||
return lambda: transport.get(path, headers=transport.master, params=params, response_type=response_type)
|
||||
return lambda: transport.get(
|
||||
path, headers=self.management_headers(transport=transport), params=params, response_type=response_type
|
||||
)
|
||||
|
||||
def model_info(self) -> list[ModelInfoEntry]:
|
||||
"""Every configured deployment with the price the proxy resolved for it
|
||||
|
|
@ -536,7 +559,7 @@ class ProxyClient:
|
|||
return unwrap(
|
||||
self.transport.get(
|
||||
"/model/info",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
params=NoBody(),
|
||||
response_type=ModelInfoResponse,
|
||||
)
|
||||
|
|
@ -546,7 +569,7 @@ class ProxyClient:
|
|||
return unwrap(
|
||||
self.transport.get(
|
||||
"/public/litellm_model_cost_map",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
params=NoBody(),
|
||||
response_type=CostMap,
|
||||
)
|
||||
|
|
@ -607,7 +630,7 @@ class ProxyClient:
|
|||
model_id = unwrap(
|
||||
self.transport.post(
|
||||
"/model/new",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=body,
|
||||
response_type=ModelNewResponse,
|
||||
)
|
||||
|
|
@ -623,7 +646,7 @@ class ProxyClient:
|
|||
|
||||
def _await_model_servable(self, model_name: str, listed_for: str | None = None) -> None:
|
||||
"""Block until every replica lists `model_name`, or fail at model_servable_timeout."""
|
||||
headers: Final = self.transport.master if listed_for is None else self.transport.bearer(listed_for)
|
||||
headers: Final = self.management_headers(listed_for)
|
||||
outcome: Final = await_servable_everywhere(
|
||||
{url: self._models_poller(transport, headers) for url, transport in self.replicas.items()},
|
||||
model_name=model_name,
|
||||
|
|
@ -666,7 +689,7 @@ class ProxyClient:
|
|||
unwrap(
|
||||
self.transport.post(
|
||||
"/model/update",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=ModelUpdateBody(
|
||||
litellm_params=litellm_params,
|
||||
model_info=ModelInfoBody(id=model_id),
|
||||
|
|
@ -678,7 +701,7 @@ class ProxyClient:
|
|||
def delete_model(self, model_id: str) -> None:
|
||||
result = self.transport.post(
|
||||
"/model/delete",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=ModelDeleteBody(id=model_id),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -747,11 +770,10 @@ class ProxyClient:
|
|||
f"GET {path} on {replica} still answers {self.poll_timeout}s after the delete; last read: {last}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _reader[R: BaseModel](transport: Transport, path: str, response_type: type[R]) -> ReplicaRead[Result[R]]:
|
||||
def _reader[R: BaseModel](self, transport: Transport, path: str, response_type: type[R]) -> ReplicaRead[Result[R]]:
|
||||
return lambda request_timeout: transport.get(
|
||||
path,
|
||||
headers=transport.master,
|
||||
headers=self.management_headers(transport=transport),
|
||||
params=NoBody(),
|
||||
response_type=response_type,
|
||||
timeout=request_timeout,
|
||||
|
|
@ -763,7 +785,7 @@ class ProxyClient:
|
|||
return unwrap(
|
||||
self.transport.post(
|
||||
"/v1/mcp/toolset",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=body,
|
||||
response_type=ToolsetRow,
|
||||
)
|
||||
|
|
@ -775,7 +797,7 @@ class ProxyClient:
|
|||
return unwrap(
|
||||
self.transport.put(
|
||||
"/v1/mcp/toolset",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=body,
|
||||
response_type=ToolsetRow,
|
||||
)
|
||||
|
|
@ -786,7 +808,7 @@ class ProxyClient:
|
|||
can unwrap it while a deferred teardown can ignore an already-deleted row."""
|
||||
return self.transport.delete(
|
||||
f"/v1/mcp/toolset/{toolset_id}",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=NoBody(),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -795,7 +817,7 @@ class ProxyClient:
|
|||
unwrap(
|
||||
self.transport.post(
|
||||
"/credentials",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=body,
|
||||
response_type=CredentialCreateResponse,
|
||||
)
|
||||
|
|
@ -804,7 +826,7 @@ class ProxyClient:
|
|||
def delete_credential(self, credential_name: str) -> None:
|
||||
result = self.transport.delete(
|
||||
f"/credentials/{credential_name}",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=NoBody(),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -815,7 +837,7 @@ class ProxyClient:
|
|||
return unwrap(
|
||||
self.transport.post(
|
||||
"/team/new",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=body,
|
||||
response_type=TeamNewResponse,
|
||||
)
|
||||
|
|
@ -824,7 +846,7 @@ class ProxyClient:
|
|||
def delete_team(self, team_id: str) -> None:
|
||||
result = self.transport.post(
|
||||
"/team/delete",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=TeamDeleteBody(team_ids=[team_id]),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
|
@ -836,7 +858,7 @@ class ProxyClient:
|
|||
a user the proxy only upserts after a successful auth."""
|
||||
result = self.transport.post(
|
||||
"/user/delete",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
json=UserDeleteBody(user_ids=[user_id]),
|
||||
response_type=UserDeleteResponse,
|
||||
)
|
||||
|
|
@ -909,7 +931,7 @@ class ProxyClient:
|
|||
def spend_logs(self, params: SpendLogsParams) -> list[SpendLogRow]:
|
||||
result = self.transport.get(
|
||||
"/spend/logs",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
params=params,
|
||||
response_type=SpendLogs,
|
||||
)
|
||||
|
|
@ -924,7 +946,7 @@ class ProxyClient:
|
|||
return unwrap(
|
||||
self.transport.get(
|
||||
"/spend/logs/v2",
|
||||
headers=self.transport.master,
|
||||
headers=self.management_headers(),
|
||||
params=SpendLogsPageParams(
|
||||
start_date=start.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
end_date=end.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
|
|
@ -977,7 +999,7 @@ class ProxyClient:
|
|||
# ---- route probe ----------------------------------------------------
|
||||
|
||||
def probe(self, path: str, *, params: NoBody) -> ProbeResult:
|
||||
return self.transport.probe(path, params=params)
|
||||
return self.transport.probe(path, params=params, headers=self.management_headers())
|
||||
|
||||
|
||||
def build_proxy_client(
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from e2e_http import (
|
|||
request_with_retry,
|
||||
streaming_outcome,
|
||||
wire_body,
|
||||
without_retries,
|
||||
)
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
|
|
@ -56,6 +57,15 @@ def _issue_from(responses: Sequence[FakeResponse]) -> Callable[[], FakeResponse]
|
|||
|
||||
|
||||
class TestTransientRetryPolicy:
|
||||
def test_qualification_disables_retries_and_restores_the_default(self) -> None:
|
||||
responses: Final = (FakeResponse(529), FakeResponse(200))
|
||||
sleep: Final = SleepRecorder()
|
||||
with without_retries():
|
||||
assert request_with_retry(_issue_from(responses), sleep=sleep) is responses[0]
|
||||
assert sleep.delays == ()
|
||||
assert request_with_retry(_issue_from(responses), sleep=sleep) is responses[1]
|
||||
assert sleep.delays == (0.5,)
|
||||
|
||||
def test_transient_set_is_only_statuses_the_proxy_cannot_emit(self) -> None:
|
||||
assert TRANSIENT_STATUSES == frozenset({529})
|
||||
assert 429 not in TRANSIENT_STATUSES
|
||||
|
|
|
|||
|
|
@ -4,12 +4,20 @@ these carry no `e2e` marker and run everywhere."""
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import signal
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from builtins import ExceptionGroup
|
||||
from collections.abc import Callable, Generator
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from threading import Thread
|
||||
from typing import Final
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from e2e_http import ExternalWrite
|
||||
|
|
@ -18,6 +26,8 @@ from idp import (
|
|||
KEYCLOAK_ADMIN_USER_ENV,
|
||||
KEYCLOAK_REALM_ENV,
|
||||
KEYCLOAK_URL_ENV,
|
||||
BrowserClientBody,
|
||||
Discovery,
|
||||
Keycloak,
|
||||
PasswordCredential,
|
||||
UserCreateBody,
|
||||
|
|
@ -60,24 +70,48 @@ def _idp_server(
|
|||
) -> Generator[tuple[Keycloak, SimpleQueue[str]]]:
|
||||
"""Exercise provisioning failures through the same HTTP transport as live tests."""
|
||||
deletions: SimpleQueue[str] = SimpleQueue()
|
||||
clients: SimpleQueue[BrowserClientBody] = SimpleQueue()
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def log_message(self, format: str, *args: object) -> None:
|
||||
pass
|
||||
|
||||
def do_POST(self) -> None:
|
||||
self.rfile.read(int(self.headers.get("Content-Length", "0")))
|
||||
body: Final = self.rfile.read(int(self.headers.get("Content-Length", "0")))
|
||||
if self.path.endswith("/token"):
|
||||
self.send_response(admin_status)
|
||||
self.end_headers()
|
||||
self.wfile.write(b'{"access_token":"synthetic-harness-token"}')
|
||||
else:
|
||||
if self.path.endswith("/clients"):
|
||||
clients.put(BrowserClientBody.model_validate_json(body))
|
||||
self.send_response(user_status if self.path.endswith("/users") else 201)
|
||||
self.send_header("Location", f"{self.path}/resource-1")
|
||||
self.end_headers()
|
||||
if user_status != 201 and self.path.endswith("/users"):
|
||||
self.wfile.write(b"injected create failure")
|
||||
|
||||
def do_GET(self) -> None:
|
||||
self.send_response(200)
|
||||
self.end_headers()
|
||||
if "/clients/" in self.path:
|
||||
client: Final = clients.get_nowait()
|
||||
clients.put(client)
|
||||
self.wfile.write(client.model_dump_json(by_alias=True).encode())
|
||||
else:
|
||||
issuer: Final = f"http://127.0.0.1:{server.server_port}/realms/test"
|
||||
self.wfile.write(
|
||||
Discovery(
|
||||
issuer=issuer,
|
||||
authorization_endpoint=f"{issuer}/auth",
|
||||
token_endpoint=f"{issuer}/token",
|
||||
userinfo_endpoint=f"{issuer}/userinfo",
|
||||
jwks_uri=f"{issuer}/certs",
|
||||
)
|
||||
.model_dump_json()
|
||||
.encode()
|
||||
)
|
||||
|
||||
def do_DELETE(self) -> None:
|
||||
deletions.put(self.path)
|
||||
self.send_response(delete_status)
|
||||
|
|
@ -115,6 +149,68 @@ def test_partial_provisioning_removes_the_group_when_user_creation_fails() -> No
|
|||
assert deletions.empty()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("exit_mode", "ignore_termination"), (("normal", False), ("parent", False), ("group", False), ("parent", True))
|
||||
)
|
||||
def test_oidc_launcher_removes_client_on_exit_and_termination(
|
||||
tmp_path: Path, exit_mode: Literal["normal", "parent", "group"], ignore_termination: bool
|
||||
) -> None:
|
||||
ready: Final = tmp_path / "ready"
|
||||
descendant_command: Final = (
|
||||
"import signal,socket,time; from pathlib import Path; "
|
||||
+ ("signal.signal(signal.SIGTERM, signal.SIG_IGN); " if ignore_termination else "")
|
||||
+ "listener=socket.socket(); listener.bind(('127.0.0.1',0)); listener.listen(); "
|
||||
f"Path({str(ready)!r}).write_text(str(listener.getsockname()[1])); time.sleep(120)"
|
||||
)
|
||||
child_command: Final = (
|
||||
"import os,subprocess,sys,time; from pathlib import Path; "
|
||||
'assert os.environ["GENERIC_CLIENT_SECRET"]; '
|
||||
'assert os.environ["GENERIC_CLIENT_USE_PKCE"] == "true"; '
|
||||
f"subprocess.Popen([sys.executable, '-c', {descendant_command!r}]); "
|
||||
f"ready=Path({str(ready)!r})\n"
|
||||
"while not ready.exists(): time.sleep(0.05)\n"
|
||||
+ ("raise SystemExit(7)" if exit_mode == "normal" else "time.sleep(120)")
|
||||
)
|
||||
with _idp_server() as (idp, deletions):
|
||||
with subprocess.Popen(
|
||||
[
|
||||
sys.executable,
|
||||
str(Path(__file__).with_name("idp.py")),
|
||||
"http://127.0.0.1:9999",
|
||||
sys.executable,
|
||||
"-c",
|
||||
child_command,
|
||||
],
|
||||
env={
|
||||
**os.environ,
|
||||
KEYCLOAK_URL_ENV: idp.base_url,
|
||||
KEYCLOAK_REALM_ENV: idp.realm,
|
||||
KEYCLOAK_ADMIN_USER_ENV: idp.admin_username,
|
||||
KEYCLOAK_ADMIN_PASSWORD_ENV: idp.admin_password,
|
||||
},
|
||||
start_new_session=True,
|
||||
) as process:
|
||||
try:
|
||||
deadline: Final = time.monotonic() + 15
|
||||
while not ready.exists() and time.monotonic() < deadline and process.poll() is None:
|
||||
time.sleep(0.05)
|
||||
assert ready.exists(), "OIDC child did not start"
|
||||
if exit_mode == "parent":
|
||||
process.terminate()
|
||||
elif exit_mode == "group":
|
||||
os.killpg(process.pid, signal.SIGTERM)
|
||||
assert process.wait(timeout=15) == (7 if exit_mode == "normal" else 143)
|
||||
with socket.socket() as connection:
|
||||
connection.settimeout(1)
|
||||
assert connection.connect_ex(("127.0.0.1", int(ready.read_text()))) != 0
|
||||
finally:
|
||||
if process.poll() is None:
|
||||
os.killpg(process.pid, signal.SIGKILL)
|
||||
process.wait(timeout=5)
|
||||
assert deletions.get(timeout=5) == "/admin/realms/test/clients/resource-1"
|
||||
assert deletions.empty()
|
||||
|
||||
|
||||
def test_successful_provisioning_cleans_up_user_before_group() -> None:
|
||||
with _idp_server() as (idp, deletions):
|
||||
with ExitStack() as cleanup:
|
||||
|
|
@ -134,6 +230,43 @@ def test_cleanup_failure_is_visible() -> None:
|
|||
idp.delete_group("group")
|
||||
|
||||
|
||||
def test_strict_cleanup_reports_each_failure_and_continues() -> None:
|
||||
from lifecycle import ResourceManager
|
||||
from proxy_client import build_proxy_client
|
||||
|
||||
with _idp_server(delete_status=500) as (idp, deletions):
|
||||
resources: Final = ResourceManager(client=build_proxy_client(), strict_cleanup=True)
|
||||
strict: Final = idp.with_strict_cleanup()
|
||||
resources.defer(lambda: strict.delete_group("group"))
|
||||
resources.defer(lambda: strict.delete_user("user"))
|
||||
with pytest.raises(ExceptionGroup, match="Resource cleanup failed") as error:
|
||||
resources.teardown()
|
||||
assert len(error.value.exceptions) == 2
|
||||
assert deletions.get_nowait() == "/admin/realms/test/users/user"
|
||||
assert deletions.get_nowait() == "/admin/realms/test/groups/group"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("groups", ((), ("one",), ("one", "two")))
|
||||
def test_provisioning_records_zero_one_or_multiple_groups(groups: tuple[str, ...]) -> None:
|
||||
with _idp_server() as (idp, deletions):
|
||||
with ExitStack() as cleanup:
|
||||
|
||||
def defer(callback: Callable[[], object]) -> None:
|
||||
cleanup.callback(callback)
|
||||
|
||||
identity: Final = idp.provision_groups(
|
||||
marker="memberships",
|
||||
groups=groups,
|
||||
defer=defer,
|
||||
)
|
||||
assert identity.groups == groups
|
||||
assert len(identity.group_ids) == len(groups)
|
||||
assert deletions.get_nowait() == "/admin/realms/test/users/resource-1"
|
||||
for _ in groups:
|
||||
assert deletions.get_nowait() == "/admin/realms/test/groups/resource-1"
|
||||
assert deletions.empty()
|
||||
|
||||
|
||||
def test_expired_admin_credentials_do_not_abort_remaining_cleanups() -> None:
|
||||
with _idp_server(admin_status=401) as (idp, _):
|
||||
cleanup: Final = ExitStack()
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ from concurrent.futures import ThreadPoolExecutor
|
|||
from contextlib import contextmanager
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -81,9 +82,14 @@ def json_object(body: bytes) -> dict[str, object]:
|
|||
class _FakeProvider(ThreadingHTTPServer):
|
||||
daemon_threads = True
|
||||
|
||||
def __init__(self, bind: tuple[str, int]) -> None:
|
||||
def __init__(self, bind: tuple[str, int], *, echo_request: bool = True) -> None:
|
||||
super().__init__(bind, _FakeProviderHandler)
|
||||
self.hits: list[str] = []
|
||||
self.echo_request = echo_request
|
||||
self.requests: tuple[tuple[Mapping[str, str], bytes], ...] = ()
|
||||
|
||||
def capture_request(self, headers: Mapping[str, str], body: bytes) -> None:
|
||||
self.requests = (*self.requests, (MappingProxyType(dict(headers)), body))
|
||||
|
||||
|
||||
class _FakeProviderHandler(BaseHTTPRequestHandler):
|
||||
|
|
@ -101,8 +107,11 @@ class _FakeProviderHandler(BaseHTTPRequestHandler):
|
|||
length = int(self.headers.get("content-length") or "0")
|
||||
body = self.rfile.read(length) if length else b""
|
||||
provider.hits.append(f"{self.command} {self.path}")
|
||||
payload = json.dumps(
|
||||
provider.capture_request(dict(self.headers.items()), body)
|
||||
payload: Final = json.dumps(
|
||||
{"echo": body.decode("utf-8"), "path": self.path, "hit": len(provider.hits)}
|
||||
if provider.echo_request
|
||||
else {"ok": True}
|
||||
).encode()
|
||||
self.send_response(200)
|
||||
self.send_header("content-type", "application/json")
|
||||
|
|
@ -117,8 +126,8 @@ class _FakeProviderHandler(BaseHTTPRequestHandler):
|
|||
|
||||
|
||||
@contextmanager
|
||||
def fake_provider() -> Generator[_FakeProvider]:
|
||||
server = _FakeProvider(("127.0.0.1", 0))
|
||||
def fake_provider(*, echo_request: bool = True) -> Generator[_FakeProvider]:
|
||||
server = _FakeProvider(("127.0.0.1", 0), echo_request=echo_request)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -11,19 +11,53 @@ injected clock, so nothing here monkeypatches anything.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Mapping
|
||||
import json
|
||||
from builtins import ExceptionGroup
|
||||
from collections.abc import Callable, Generator, Iterable, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from itertools import chain, repeat
|
||||
from queue import SimpleQueue
|
||||
from threading import Thread
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast
|
||||
|
||||
import pytest
|
||||
from e2e_config import parse_replica_urls
|
||||
from e2e_http import Result, Success
|
||||
from models import KeyInfo, KeyInfoResponse, ModelListEntry, ModelsListResponse
|
||||
from e2e_http import NoBody, Result, Success, without_retries
|
||||
from idp import Keycloak
|
||||
from lifecycle import ResourceManager
|
||||
from management.jwt_actors import ActorFactory
|
||||
from management.management_client import ManagementClient
|
||||
from models import (
|
||||
ConnectionTestBody,
|
||||
CredentialCreateBody,
|
||||
KeyGenerateBody,
|
||||
KeyInfo,
|
||||
KeyInfoResponse,
|
||||
KeyUpdateBody,
|
||||
LiteLLMParamsBody,
|
||||
McpServerCreateBody,
|
||||
McpServerUpdateBody,
|
||||
ModelListEntry,
|
||||
ModelsListResponse,
|
||||
OrgNewBody,
|
||||
OrgUpdateBody,
|
||||
SpendLogsParams,
|
||||
TagNewBody,
|
||||
TeamNewBody,
|
||||
TeamUpdateBody,
|
||||
ToolsetCreateBody,
|
||||
ToolsetUpdateBody,
|
||||
UserNewBody,
|
||||
UserUpdateBody,
|
||||
)
|
||||
from proxy_client import (
|
||||
ConvergeOutcome,
|
||||
Caller,
|
||||
Converged,
|
||||
ConvergeOutcome,
|
||||
CredentialKind,
|
||||
EverywhereConverged,
|
||||
ModelsPoller,
|
||||
NeverConvergedOn,
|
||||
|
|
@ -42,6 +76,115 @@ from proxy_client import (
|
|||
)
|
||||
from transport import Transport
|
||||
|
||||
|
||||
@contextmanager
|
||||
def caller_boundary(
|
||||
status: int = 200, bodies: SimpleQueue[bytes] | None = None, *, delete_status: int | None = None
|
||||
) -> Generator[tuple[ManagementClient, SimpleQueue[str]]]:
|
||||
received: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
def log_message(self, format: str, *args: object) -> None:
|
||||
pass
|
||||
|
||||
def do_GET(self) -> None:
|
||||
received.put(self.headers.get("Authorization", ""))
|
||||
self.send_response(delete_status if self.path == "/key/delete" and delete_status is not None else status)
|
||||
self.end_headers()
|
||||
self.wfile.write(
|
||||
b'{"key":"owned","info":{"key_alias":"owned"},"data":[{"id":"owned"}],"team_id":"owned","team_info":{},"model_id":"owned"}'
|
||||
)
|
||||
|
||||
def do_POST(self) -> None:
|
||||
body: Final = self.rfile.read(int(self.headers.get("Content-Length", "0")))
|
||||
if bodies is not None:
|
||||
bodies.put(body)
|
||||
self.do_GET()
|
||||
|
||||
do_PATCH = do_POST
|
||||
do_PUT = do_POST
|
||||
do_DELETE = do_POST
|
||||
|
||||
server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
thread: Final = Thread(target=lambda: server.serve_forever(poll_interval=0.01), daemon=True)
|
||||
thread.start()
|
||||
url: Final = f"http://127.0.0.1:{server.server_port}"
|
||||
proxy: Final = build_proxy_client(
|
||||
base_url=url, control_plane_base_url=url, replica_urls=(url,), master_key="bootstrap"
|
||||
)
|
||||
try:
|
||||
yield ManagementClient(proxy=proxy, master_key="bootstrap"), received
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
class TestBoundManagementCaller:
|
||||
def test_strict_key_cleanup_accepts_missing_only_when_requested(self) -> None:
|
||||
with caller_boundary(delete_status=404) as (bootstrap, received), without_retries():
|
||||
with pytest.raises(AssertionError):
|
||||
bootstrap.delete_key_strict("owned")
|
||||
bootstrap.delete_key_strict("owned", missing_ok=True)
|
||||
assert (received.get_nowait(), received.get_nowait()) == ("Bearer bootstrap", "Bearer bootstrap")
|
||||
|
||||
def test_actor_key_cleanup_reports_failure_and_continues(self) -> None:
|
||||
with caller_boundary(delete_status=500) as (bootstrap, received), without_retries():
|
||||
resources: Final = ResourceManager(client=bootstrap.proxy, strict_cleanup=True)
|
||||
remaining: SimpleQueue[str] = SimpleQueue()
|
||||
resources.defer(lambda: remaining.put("cleaned"))
|
||||
factory: Final = ActorFactory(
|
||||
bootstrap=bootstrap,
|
||||
idp=Keycloak(base_url="http://unused.test", realm="test", admin_username="test", admin_password="test"),
|
||||
resources=resources,
|
||||
)
|
||||
assert factory.key().key == "owned"
|
||||
with pytest.raises(ExceptionGroup, match="Resource cleanup failed") as failure:
|
||||
resources.teardown()
|
||||
assert len(failure.value.exceptions) == 1
|
||||
assert remaining.get_nowait() == "cleaned"
|
||||
assert (received.get_nowait(), received.get_nowait()) == ("Bearer bootstrap", "Bearer bootstrap")
|
||||
|
||||
@pytest.mark.parametrize("kind", ("direct_jwt", "virtual_key", "dashboard_session"))
|
||||
def test_direct_delegated_and_replica_reads_keep_the_bound_caller(self, kind: CredentialKind) -> None:
|
||||
with caller_boundary() as (bootstrap, received):
|
||||
caller: Final = Caller(credential="synthetic-caller", kind=kind, role="internal_user", tenant="tenant-a")
|
||||
bound: Final = bootstrap.with_caller(caller)
|
||||
bound.update_key(KeyUpdateBody(key="owned", key_alias="updated"))
|
||||
bound.proxy.key_info("owned")
|
||||
bound.proxy.read_back_everywhere(
|
||||
"/key/info",
|
||||
params=KeyUpdateBody(key="owned"),
|
||||
response_type=KeyInfoResponse,
|
||||
converged=lambda result: isinstance(result, Success),
|
||||
)
|
||||
bound.proxy.read_body_back_everywhere(
|
||||
"/key/info", KeyInfoResponse, settled=lambda result: result.info.key_alias == "owned"
|
||||
)
|
||||
assert tuple(received.get_nowait() for _ in range(4)) == ("Bearer synthetic-caller",) * 4
|
||||
assert received.empty()
|
||||
bootstrap.proxy.key_info("owned")
|
||||
assert received.get_nowait() == "Bearer bootstrap"
|
||||
|
||||
def test_explicit_override_wins_without_rebinding_or_changing_master(self) -> None:
|
||||
with caller_boundary() as (bootstrap, received):
|
||||
bound: Final = bootstrap.with_caller(Caller(credential="bound", kind="direct_jwt", role="internal_user"))
|
||||
bound.update_key(KeyUpdateBody(key="owned"), caller_key="override")
|
||||
bound.proxy.key_info("owned")
|
||||
assert received.get_nowait() == "Bearer override"
|
||||
assert received.get_nowait() == "Bearer bound"
|
||||
assert bound.master_key == "bootstrap"
|
||||
|
||||
def test_credentials_are_absent_from_binding_and_header_diagnostics(self) -> None:
|
||||
with caller_boundary() as (bootstrap, _):
|
||||
caller: Final = Caller(credential="private-value", kind="direct_jwt", role="internal_user")
|
||||
bound: Final = bootstrap.with_caller(caller)
|
||||
assert "private-value" not in repr(caller)
|
||||
assert "private-value" not in repr(bound)
|
||||
assert "private-value" not in repr(bound.proxy.management_headers())
|
||||
assert "bootstrap" not in repr(bound)
|
||||
|
||||
|
||||
MODEL: Final = "gpt-under-test"
|
||||
_NO_TRANSPORTS: Final = cast(Transport, None)
|
||||
TIMEOUT: Final = 10.0
|
||||
|
|
@ -275,3 +418,166 @@ class TestReplicasFor:
|
|||
client: Final = ProxyClient(transport=_NO_TRANSPORTS, replicas={}, control_replicas={})
|
||||
with pytest.raises(AssertionError, match="no replica is configured"):
|
||||
_ = client.replicas_for("/v1/models")
|
||||
|
||||
|
||||
MANAGEMENT_OPERATIONS: Final[tuple[tuple[str, Callable[[ManagementClient], object]], ...]] = (
|
||||
("generate_key", lambda c: c.generate_key(KeyGenerateBody())),
|
||||
("llm_only_key", lambda c: c.llm_only_key()),
|
||||
("update_key", lambda c: c.update_key(KeyUpdateBody(key="owned"))),
|
||||
("update_key_models", lambda c: c.update_key_models("owned", [])),
|
||||
("key_info", lambda c: c.key_info_as("owned")),
|
||||
("delete_key_strict", lambda c: c.delete_key_strict("owned")),
|
||||
("delete_model_strict", lambda c: c.delete_model_strict("owned")),
|
||||
(
|
||||
"connection_test",
|
||||
lambda c: c.connection_test(
|
||||
ConnectionTestBody(litellm_params=LiteLLMParamsBody(model="synthetic"), mode="chat")
|
||||
),
|
||||
),
|
||||
("block_key", lambda c: c.block_key("owned")),
|
||||
("regenerate_key", lambda c: c.regenerate_key("owned")),
|
||||
("reset_key_spend", lambda c: c.reset_key_spend("owned", 0)),
|
||||
("key_list", lambda c: c.key_list("owned")),
|
||||
("key_alias_count", lambda c: c.key_alias_count("owned")),
|
||||
("create_team", lambda c: c.create_team(TeamNewBody(team_alias="owned"))),
|
||||
("update_team", lambda c: c.update_team(TeamUpdateBody(team_id="owned", team_alias="updated"))),
|
||||
("delete_team", lambda c: c.delete_team("owned")),
|
||||
("team_info", lambda c: c.team_info("owned")),
|
||||
("team_list_ids", lambda c: c.team_list_ids()),
|
||||
("team_info_status", lambda c: c.team_info_status("owned")),
|
||||
("add_team_member", lambda c: c.add_team_member("owned", "user")),
|
||||
("delete_team_member", lambda c: c.delete_team_member("owned", "user")),
|
||||
("create_user", lambda c: c.create_user(UserNewBody(user_email="actor@example.com", user_role="internal_user"))),
|
||||
("create_customer", lambda c: c.create_customer("owned")),
|
||||
("customer_info", lambda c: c.customer_info("owned")),
|
||||
("delete_customer", lambda c: c.delete_customer("owned")),
|
||||
("update_user", lambda c: c.update_user(UserUpdateBody(user_id="owned", user_role="internal_user"))),
|
||||
("delete_user", lambda c: c.delete_user("owned")),
|
||||
("delete_user_strict", lambda c: c.delete_user_strict("owned")),
|
||||
("user_info", lambda c: c.user_info("owned")),
|
||||
("user_count", lambda c: c.user_count("owned")),
|
||||
("user_list_ids", lambda c: c.user_list_ids("owned")),
|
||||
("create_org", lambda c: c.create_org(OrgNewBody(organization_alias="owned"))),
|
||||
("update_org", lambda c: c.update_org(OrgUpdateBody(organization_id="owned", organization_alias="updated"))),
|
||||
("delete_org", lambda c: c.delete_org("owned")),
|
||||
("org_info", lambda c: c.org_info("owned")),
|
||||
("org_info_status", lambda c: c.org_info_status("owned")),
|
||||
("create_tag", lambda c: c.create_tag(TagNewBody(name="owned"))),
|
||||
("delete_tag", lambda c: c.delete_tag("owned")),
|
||||
("tag_list", lambda c: c.tag_list()),
|
||||
("create_mcp_server", lambda c: c.create_mcp_server(McpServerCreateBody(alias="owned", url="http://example.test"))),
|
||||
("update_mcp_server", lambda c: c.update_mcp_server(McpServerUpdateBody(server_id="owned", alias=None))),
|
||||
("delete_mcp_server", lambda c: c.delete_mcp_server("owned")),
|
||||
("proxy.generate_key", lambda c: c.proxy.generate_key(KeyGenerateBody())),
|
||||
("proxy.delete_key", lambda c: c.proxy.delete_key("owned")),
|
||||
("proxy.delete_customers", lambda c: c.proxy.delete_customers(["owned"])),
|
||||
("proxy.key_info", lambda c: c.proxy.key_info("owned")),
|
||||
("proxy.memory_summary", lambda c: c.proxy.memory_summary_everywhere()),
|
||||
("proxy.model_info", lambda c: c.proxy.model_info()),
|
||||
("proxy.model_cost_map", lambda c: c.proxy.model_cost_map()),
|
||||
("proxy.create_model", lambda c: c.proxy.create_model("owned", LiteLLMParamsBody(model="synthetic"))),
|
||||
("proxy.update_model", lambda c: c.proxy.update_model("owned", LiteLLMParamsBody(model="synthetic"))),
|
||||
("proxy.delete_model", lambda c: c.proxy.delete_model("owned")),
|
||||
("proxy.create_toolset", lambda c: c.proxy.create_toolset(ToolsetCreateBody(toolset_name="owned", tools=[]))),
|
||||
("proxy.update_toolset", lambda c: c.proxy.update_toolset(ToolsetUpdateBody(toolset_id="owned", description=None))),
|
||||
("proxy.delete_toolset", lambda c: c.proxy.delete_toolset("owned")),
|
||||
(
|
||||
"proxy.create_credential",
|
||||
lambda c: c.proxy.create_credential(CredentialCreateBody(credential_name="owned", credential_values={})),
|
||||
),
|
||||
("proxy.delete_credential", lambda c: c.proxy.delete_credential("owned")),
|
||||
("proxy.create_team", lambda c: c.proxy.create_team(TeamNewBody(team_alias="owned"))),
|
||||
("proxy.delete_team", lambda c: c.proxy.delete_team("owned")),
|
||||
("proxy.delete_user", lambda c: c.proxy.delete_user("owned")),
|
||||
("proxy.spend_logs", lambda c: c.proxy.spend_logs(SpendLogsParams(api_key="owned"))),
|
||||
("proxy.probe", lambda c: c.proxy.probe("/user/info", params=NoBody())),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("name", "operation"), MANAGEMENT_OPERATIONS, ids=tuple(name for name, _ in MANAGEMENT_OPERATIONS)
|
||||
)
|
||||
@pytest.mark.parametrize("kind", ("master", "direct_jwt", "virtual_key", "dashboard_session"))
|
||||
def test_management_operations_send_the_selected_credential(
|
||||
name: str,
|
||||
operation: Callable[[ManagementClient], object],
|
||||
kind: CredentialKind,
|
||||
) -> None:
|
||||
with caller_boundary(status=401) as (bootstrap, received), without_retries():
|
||||
client: Final = (
|
||||
bootstrap
|
||||
if kind == "master"
|
||||
else bootstrap.with_caller(Caller(credential=f"synthetic-{kind}", kind=kind, role="internal_user"))
|
||||
)
|
||||
try:
|
||||
operation(client)
|
||||
except AssertionError:
|
||||
pass
|
||||
expected: Final = "Bearer bootstrap" if kind == "master" else f"Bearer synthetic-{kind}"
|
||||
assert received.get_nowait() == expected, name
|
||||
assert received.empty(), "an unauthorized request must not be retried"
|
||||
|
||||
|
||||
class TestSplitCallerPropagation:
|
||||
def test_control_and_data_replica_readers_keep_the_caller(self) -> None:
|
||||
with caller_boundary() as (data, data_headers), caller_boundary() as (control, control_headers):
|
||||
data_url: Final = next(iter(data.proxy.replicas))
|
||||
control_url: Final = next(iter(control.proxy.replicas))
|
||||
proxy: Final = build_proxy_client(
|
||||
base_url=data_url,
|
||||
control_plane_base_url=control_url,
|
||||
replica_urls=(data_url,),
|
||||
master_key="bootstrap",
|
||||
).with_caller(Caller(credential="tenant-token", kind="direct_jwt", role="team_member"))
|
||||
proxy.key_info("owned")
|
||||
proxy.read_body_back_everywhere(
|
||||
"/key/info", KeyInfoResponse, settled=lambda info: info.info.key_alias == "owned"
|
||||
)
|
||||
proxy.read_back_everywhere(
|
||||
"/key/info",
|
||||
params=NoBody(),
|
||||
response_type=KeyInfoResponse,
|
||||
converged=lambda result: isinstance(result, Success),
|
||||
)
|
||||
assert control_headers.get_nowait() == "Bearer tenant-token"
|
||||
assert control_headers.get_nowait() == "Bearer tenant-token"
|
||||
assert data_headers.get_nowait() == "Bearer tenant-token"
|
||||
assert control_headers.empty() and data_headers.empty()
|
||||
|
||||
def test_successful_team_and_model_polling_uses_the_bound_caller(self) -> None:
|
||||
with caller_boundary() as (bootstrap, received):
|
||||
bound: Final = bootstrap.with_caller(Caller(credential="caller", kind="direct_jwt", role="proxy_admin"))
|
||||
bound.create_team(TeamNewBody(team_alias="owned"))
|
||||
bound.proxy.create_model("owned", LiteLLMParamsBody(model="synthetic"))
|
||||
assert tuple(received.get_nowait() for _ in range(4)) == ("Bearer caller",) * 4
|
||||
assert received.empty()
|
||||
|
||||
def test_expired_shaped_token_is_sent_once_without_renewal(self) -> None:
|
||||
with caller_boundary(status=401) as (bootstrap, received):
|
||||
bound: Final = bootstrap.with_caller(
|
||||
Caller(credential="expired.payload.signature", kind="direct_jwt", role="internal_user")
|
||||
)
|
||||
result: Final = bound.key_info_as("owned")
|
||||
assert not isinstance(result, Success)
|
||||
assert received.get_nowait() == "Bearer expired.payload.signature"
|
||||
assert received.empty()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("operation", ("server", "toolset"))
|
||||
def test_partial_updates_preserve_explicit_null_at_the_http_boundary(operation: str) -> None:
|
||||
bodies: Final[SimpleQueue[bytes]] = SimpleQueue()
|
||||
with caller_boundary(status=401, bodies=bodies) as (bootstrap, _):
|
||||
try:
|
||||
if operation == "server":
|
||||
bootstrap.update_mcp_server(McpServerUpdateBody(server_id="owned", alias=None))
|
||||
else:
|
||||
bootstrap.proxy.update_toolset(ToolsetUpdateBody(toolset_id="owned", description=None))
|
||||
except AssertionError:
|
||||
pass
|
||||
expected: Final = (
|
||||
{"server_id": "owned", "alias": None}
|
||||
if operation == "server"
|
||||
else {"toolset_id": "owned", "description": None}
|
||||
)
|
||||
assert json.loads(bodies.get_nowait()) == expected
|
||||
assert bodies.empty()
|
||||
|
|
|
|||
|
|
@ -7,11 +7,9 @@ client touches requests.* or builds raw dicts; they pass pydantic models here.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
import e2e_http
|
||||
from e2e_http import (
|
||||
URL,
|
||||
|
|
@ -21,6 +19,7 @@ from e2e_http import (
|
|||
Result,
|
||||
StreamingResponse,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class Transport(Protocol):
|
||||
|
|
@ -85,7 +84,7 @@ class Transport(Protocol):
|
|||
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
|
||||
) -> Result[R]: ...
|
||||
|
||||
def probe(self, path: str, *, params: BaseModel) -> ProbeResult: ...
|
||||
def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult: ...
|
||||
|
||||
def upload[R: BaseModel](
|
||||
self,
|
||||
|
|
@ -113,7 +112,7 @@ class Transport(Protocol):
|
|||
@dataclass(frozen=True, slots=True)
|
||||
class HttpTransport:
|
||||
base_url: str
|
||||
master_key: str
|
||||
master_key: str = field(repr=False)
|
||||
request_timeout: float = 60.0
|
||||
|
||||
def _url(self, path: str) -> URL:
|
||||
|
|
@ -245,10 +244,10 @@ class HttpTransport:
|
|||
timeout=self.request_timeout,
|
||||
)
|
||||
|
||||
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
|
||||
def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult:
|
||||
return e2e_http.probe(
|
||||
self._url(path),
|
||||
headers=self.master,
|
||||
headers=self.master if headers is None else headers,
|
||||
params=params,
|
||||
timeout=self.request_timeout,
|
||||
)
|
||||
|
|
@ -434,8 +433,8 @@ class SplitTransport:
|
|||
path, headers=headers, json=json, params=params, stream=stream
|
||||
)
|
||||
|
||||
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
|
||||
return self._route(path).probe(path, params=params)
|
||||
def probe(self, path: str, *, params: BaseModel, headers: BaseModel | None = None) -> ProbeResult:
|
||||
return self._route(path).probe(path, params=params, headers=headers)
|
||||
|
||||
def upload[R: BaseModel](
|
||||
self,
|
||||
|
|
|
|||
30
tests/e2e/ui/oidcSetup.ts
Normal file
30
tests/e2e/ui/oidcSetup.ts
Normal file
|
|
@ -0,0 +1,30 @@
|
|||
import { chromium, expect } from "@playwright/test";
|
||||
import * as fs from "fs";
|
||||
import * as path from "path";
|
||||
|
||||
export default async function oidcSetup() {
|
||||
const baseURL = process.env.E2E_OIDC_UI_URL;
|
||||
const issuer = process.env.JWT_ISSUER;
|
||||
const username = process.env.E2E_OIDC_USERNAME;
|
||||
const password = process.env.E2E_OIDC_PASSWORD;
|
||||
if (!baseURL || !issuer || !username || !password) {
|
||||
throw new Error("The OIDC setup requires a running stack, issuer, and provisioned actor credentials");
|
||||
}
|
||||
const artifactDir = process.env.E2E_UI_ARTIFACT_DIR || ".";
|
||||
fs.mkdirSync(artifactDir, { recursive: true });
|
||||
const browser = await chromium.launch();
|
||||
try {
|
||||
const page = await browser.newPage();
|
||||
await page.goto(`${baseURL.replace(/\/$/, "")}/sso/key/generate`);
|
||||
await expect(page).toHaveURL(new RegExp(`^${issuer.replace(/[.*+?^${}()|[\]\\]/g, "\\$&")}/`));
|
||||
await page.getByLabel("Username or email").fill(username);
|
||||
await page.getByLabel("Password", { exact: true }).fill(password);
|
||||
await page.getByRole("button", { name: "Sign In", exact: true }).click();
|
||||
await page.waitForURL((url) => url.origin === new URL(baseURL).origin && url.pathname.startsWith("/ui"));
|
||||
const statePath = path.join(artifactDir, "oidc.storageState.json");
|
||||
await page.context().storageState({ path: statePath });
|
||||
fs.chmodSync(statePath, 0o600);
|
||||
} finally {
|
||||
await browser.close();
|
||||
}
|
||||
}
|
||||
22
tests/e2e/ui/playwright.oidc.config.ts
Normal file
22
tests/e2e/ui/playwright.oidc.config.ts
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
import { defineConfig, devices } from "@playwright/test";
|
||||
import * as path from "path";
|
||||
|
||||
const baseURL = process.env.E2E_OIDC_UI_URL;
|
||||
if (!baseURL) throw new Error("E2E_OIDC_UI_URL must point to the running OIDC stack");
|
||||
|
||||
export default defineConfig({
|
||||
testDir: ".",
|
||||
testMatch: "oidc/**/*.spec.ts",
|
||||
retries: 0,
|
||||
workers: 1,
|
||||
outputDir: path.join(process.env.E2E_UI_ARTIFACT_DIR || ".", "oidc", "test-results"),
|
||||
globalSetup: require.resolve("./oidcSetup"),
|
||||
use: {
|
||||
...devices["Desktop Chrome"],
|
||||
baseURL,
|
||||
storageState: path.join(process.env.E2E_UI_ARTIFACT_DIR || ".", "oidc.storageState.json"),
|
||||
trace: "off",
|
||||
screenshot: "off",
|
||||
video: "off",
|
||||
},
|
||||
});
|
||||
|
|
@ -695,6 +695,38 @@ def test_vertex_cost_and_usage_aggregation(monkeypatch):
|
|||
assert result.failed_requests == 0
|
||||
|
||||
|
||||
def test_vertex_batch_usage_preserves_modality_token_details(monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"vertex_ai/gemini-embedding-2",
|
||||
{
|
||||
"input_cost_per_token_batches": 1e-7,
|
||||
"input_cost_per_audio_token_batches": 3.25e-6,
|
||||
"input_cost_per_image_token_batches": 2.25e-7,
|
||||
"input_cost_per_video_token_batches": 6e-6,
|
||||
},
|
||||
)
|
||||
responses = [
|
||||
{
|
||||
"response": {
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 84,
|
||||
"candidatesTokenCount": 0,
|
||||
"totalTokenCount": 84,
|
||||
"promptTokensDetails": [
|
||||
{"modality": "AUDIO", "tokenCount": 64},
|
||||
{"modality": "TEXT", "tokenCount": 20},
|
||||
],
|
||||
}
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
result = bu.calculate_vertex_ai_batch_cost_and_usage(responses, "gemini-embedding-2")
|
||||
|
||||
assert result.prompt_cost == pytest.approx(64 * 3.25e-6 + 20 * 1e-7)
|
||||
|
||||
|
||||
def test_vertex_cost_skips_none_response_body(monkeypatch):
|
||||
import litellm.cost_calculator as cc
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,8 @@ never rewrite. It is consumed by compress() and by the Headroom guardrail, so
|
|||
the two agree on what "never compress this" means.
|
||||
"""
|
||||
|
||||
from litellm.compression.compress import get_protected_indices
|
||||
from litellm.compression.compress import compress, get_protected_indices
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
def test_protects_system_last_user_and_last_assistant():
|
||||
|
|
@ -53,3 +54,94 @@ def test_every_system_row_is_protected():
|
|||
def test_no_user_or_assistant_rows():
|
||||
assert sorted(get_protected_indices([{"role": "system", "content": "sys"}])) == [0]
|
||||
assert get_protected_indices([]) == ()
|
||||
|
||||
|
||||
def test_mid_history_cache_control_part_is_protected():
|
||||
# A large cached tool result from a few turns back, not the last user or
|
||||
# last assistant row -- exactly the row a provider prompt-cache pins to
|
||||
# exact bytes. Rewriting it (even leaving the marker on) changes those
|
||||
# bytes and turns the next request's cache read into a cache write.
|
||||
messages = [
|
||||
{"role": "user", "content": "old question"},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "a large cached tool result", "cache_control": {"type": "ephemeral"}},
|
||||
],
|
||||
},
|
||||
{"role": "assistant", "content": "ack"},
|
||||
{"role": "user", "content": "live instruction"},
|
||||
]
|
||||
|
||||
# index 3 = last assistant, index 4 = last user (both protected by role
|
||||
# regardless), index 2 = the cache_control-marked row itself.
|
||||
assert sorted(get_protected_indices(messages)) == [2, 3, 4]
|
||||
|
||||
|
||||
def test_cache_control_directly_on_message_is_protected():
|
||||
messages = [
|
||||
{"role": "user", "content": "old question", "cache_control": {"type": "ephemeral"}},
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
{"role": "user", "content": "live instruction"},
|
||||
]
|
||||
|
||||
assert sorted(get_protected_indices(messages)) == [0, 1, 2]
|
||||
|
||||
|
||||
def test_cache_control_protection_does_not_duplicate_already_protected_rows():
|
||||
# The last user row is already protected by role; marking it too must not
|
||||
# produce a duplicate index.
|
||||
messages = [
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "user", "content": "live", "cache_control": {"type": "ephemeral"}},
|
||||
]
|
||||
|
||||
protected = get_protected_indices(messages)
|
||||
|
||||
assert sorted(protected) == [0, 1]
|
||||
assert len(protected) == len(set(protected))
|
||||
|
||||
|
||||
def test_content_that_is_not_a_list_of_mappings_is_not_treated_as_cache_control():
|
||||
# Defensive: a plain string content, or a list of non-dict items, must not
|
||||
# raise or be misread as carrying a breakpoint.
|
||||
messages = [
|
||||
{"role": "assistant", "content": "plain string content"},
|
||||
{"role": "user", "content": ["not", "a", "dict", "list"]},
|
||||
{"role": "user", "content": "live instruction"},
|
||||
]
|
||||
|
||||
assert sorted(get_protected_indices(messages)) == [0, 2]
|
||||
|
||||
|
||||
def test_compress_keeps_part_level_cache_control_row_verbatim():
|
||||
# compress() scores text-only copies of the rows, where a part-level marker
|
||||
# is gone; protection has to read the original rows or the pinned row is stubbed.
|
||||
stale_log = {"role": "user", "content": [{"type": "text", "text": "stale log line " * 2000}]}
|
||||
pinned = {
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "cached tool result " * 2000, "cache_control": {"type": "ephemeral"}},
|
||||
],
|
||||
}
|
||||
messages = [
|
||||
stale_log,
|
||||
{"role": "assistant", "content": "old answer"},
|
||||
pinned,
|
||||
{"role": "assistant", "content": "ack"},
|
||||
{"role": "user", "content": "live instruction"},
|
||||
]
|
||||
|
||||
result = compress(
|
||||
messages,
|
||||
model="gpt-4o",
|
||||
call_type=CallTypes.anthropic_messages,
|
||||
compression_trigger=1000,
|
||||
compression_target=500,
|
||||
)
|
||||
|
||||
assert len(result["messages"]) == len(messages)
|
||||
assert result["messages"][2] == pinned
|
||||
assert result["messages"][0] != stale_log
|
||||
assert len(result["cache"]) >= 1
|
||||
|
|
|
|||
|
|
@ -74,6 +74,108 @@ def test_missing_cache_read_policy_preserves_billing(prompt_tokens, read_rate, s
|
|||
assert prompt_cost == pytest.approx((prompt_tokens - 100) * billed[0] + 100 * billed[4])
|
||||
|
||||
|
||||
def test_generic_cost_per_token_prefers_audio_per_second_rate() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"key": "gemini-embedding-2",
|
||||
"max_tokens": None,
|
||||
"max_input_tokens": None,
|
||||
"max_output_tokens": None,
|
||||
"input_cost_per_token": 2e-7,
|
||||
"input_cost_per_audio_token": 6.5e-6,
|
||||
"input_cost_per_audio_per_second": 0.00016,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"mode": "embedding",
|
||||
"supported_openai_params": None,
|
||||
}
|
||||
usage = Usage(
|
||||
prompt_tokens=64,
|
||||
completion_tokens=0,
|
||||
total_tokens=64,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
audio_tokens=64,
|
||||
audio_length_seconds=2,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gemini-embedding-2",
|
||||
usage=usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_info=model_info,
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(2 * 0.00016)
|
||||
|
||||
|
||||
def test_generic_cost_per_token_prefers_image_per_image_rate() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"key": "gemini-embedding-2",
|
||||
"max_tokens": None,
|
||||
"max_input_tokens": None,
|
||||
"max_output_tokens": None,
|
||||
"input_cost_per_token": 2e-7,
|
||||
"input_cost_per_image_token": 4.5e-7,
|
||||
"input_cost_per_image": 0.00012,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"mode": "embedding",
|
||||
"supported_openai_params": None,
|
||||
}
|
||||
usage = Usage(
|
||||
prompt_tokens=258,
|
||||
completion_tokens=0,
|
||||
total_tokens=258,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
image_tokens=258,
|
||||
image_count=1,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gemini-embedding-2",
|
||||
usage=usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_info=model_info,
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(0.00012)
|
||||
|
||||
|
||||
def test_generic_cost_per_token_prefers_video_per_second_rate() -> None:
|
||||
model_info: ModelInfo = {
|
||||
"key": "gemini-embedding-2",
|
||||
"max_tokens": None,
|
||||
"max_input_tokens": None,
|
||||
"max_output_tokens": None,
|
||||
"input_cost_per_token": 2e-7,
|
||||
"input_cost_per_video_token": 1.2e-5,
|
||||
"input_cost_per_video_per_second": 0.00079,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "vertex_ai",
|
||||
"mode": "embedding",
|
||||
"supported_openai_params": None,
|
||||
}
|
||||
usage = Usage(
|
||||
prompt_tokens=516,
|
||||
completion_tokens=0,
|
||||
total_tokens=516,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
video_tokens=516,
|
||||
video_length_seconds=2,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gemini-embedding-2",
|
||||
usage=usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_info=model_info,
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(2 * 0.00079)
|
||||
|
||||
|
||||
def test_missing_cache_read_uses_off_peak_input_rate():
|
||||
from datetime import datetime, timezone
|
||||
|
||||
|
|
|
|||
|
|
@ -2481,6 +2481,74 @@ class TestAnthropicMessagesHandlerStreamingScanKey:
|
|||
assert ended_key != open_key
|
||||
|
||||
|
||||
class PerRowTextGuardrail(CustomGuardrail):
|
||||
"""Answers one redacted text per chat row it was shown, the way a guardrail
|
||||
that scans per message does, and hands back only texts."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="per-row-redactor")
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
rows = inputs.get("structured_messages") or []
|
||||
return {**inputs, "texts": [str(row.get("content")).replace("123-45-6789", "<US_SSN>") for row in rows]}
|
||||
|
||||
|
||||
class TestPerMessageTextWriteBack:
|
||||
"""Texts that no longer pair one-to-one with what the handler extracted must be
|
||||
rejected by name instead of sliding onto the wrong messages."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_per_row_over_a_system_prompt_is_applied(self):
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": "Reply with exactly the SSN you were given.",
|
||||
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
|
||||
}
|
||||
|
||||
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
|
||||
|
||||
assert data["system"] == "Reply with exactly the SSN you were given."
|
||||
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_per_row_over_a_multi_block_system_prompt_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": [
|
||||
{"type": "text", "text": "Reply with exactly the SSN you were given."},
|
||||
{"type": "text", "text": "Never apologize."},
|
||||
],
|
||||
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
|
||||
}
|
||||
original = json.loads(json.dumps(data))
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-row-redactor"
|
||||
assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched"
|
||||
assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_per_row_without_a_system_prompt_is_applied(self):
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
|
||||
}
|
||||
|
||||
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
|
||||
|
||||
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
|
||||
|
||||
|
||||
class TestAnthropicMessagesHandlerPostCallHookResponse:
|
||||
def test_openai_shaped_stream_assembly_reaches_the_hook_as_a_messages_response(self):
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
|
|
|
|||
|
|
@ -1893,6 +1893,49 @@ class TestScanOnlyToolResults:
|
|||
assert data["messages"][4]["content"] == "and then?"
|
||||
|
||||
|
||||
class ToolDroppingTextGuardrail(CustomGuardrail):
|
||||
"""Answers one text per non-tool message it saw, the way a guardrail that
|
||||
filters tool rows out before scanning does, and hands back only texts."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="tool-dropping-redactor")
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
kept = [m for m in inputs.get("structured_messages") or [] if m.get("role") != "tool"]
|
||||
return {**inputs, "texts": [str(m.get("content")).replace("POISON", "[BLOCKED]") for m in kept]}
|
||||
|
||||
|
||||
class TestPerMessageTextWriteBack:
|
||||
"""Texts that no longer pair one-to-one with what the handler extracted must be
|
||||
rejected by name instead of sliding onto the wrong messages."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fewer_texts_than_extracted_over_a_tool_message_is_rejected(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
original_messages = [
|
||||
{"role": "system", "content": "SYSTEM-PROMPT"},
|
||||
{"role": "user", "content": "fetch the page"},
|
||||
{"role": "assistant", "content": "fetching"},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "page says POISON here"},
|
||||
{"role": "user", "content": "and then?"},
|
||||
]
|
||||
data = {"messages": json.loads(json.dumps(original_messages))}
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=ToolDroppingTextGuardrail())
|
||||
|
||||
assert excinfo.value.guardrail_name == "tool-dropping-redactor"
|
||||
assert data["messages"] == original_messages, "a rejected rewrite must leave the request untouched"
|
||||
|
||||
|
||||
class TestBuildBlockSseChunks:
|
||||
"""build_block_sse_chunks turns a streaming ModifyResponseException into 200 SSE chunks"""
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ with guardrail transformations.
|
|||
import copy
|
||||
from collections.abc import Callable
|
||||
from typing import Any, List, Literal, Optional, Tuple
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import logging
|
||||
|
||||
|
|
@ -31,6 +31,7 @@ from litellm.llms.openai.responses.guardrail_translation.handler import (
|
|||
OpenAIResponsesHandler,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools
|
||||
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
|
||||
from litellm.types.llms.openai import ChatCompletionToolCallChunk
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
|
|
@ -2338,6 +2339,135 @@ def _parallel_tool_call_input() -> list:
|
|||
]
|
||||
|
||||
|
||||
SSN = "123-45-6789"
|
||||
REDACTED_SSN = "<US_SSN>"
|
||||
|
||||
|
||||
def _redacted(value: object) -> object:
|
||||
if isinstance(value, str):
|
||||
return value.replace(SSN, REDACTED_SSN)
|
||||
if isinstance(value, list):
|
||||
return [{**part, "text": _redacted(part["text"])} if "text" in part else part for part in value]
|
||||
return value
|
||||
|
||||
|
||||
def _per_message_guardrail_server(structured_messages_in_answer: bool) -> Callable[..., MagicMock]:
|
||||
"""Answers one redacted text per chat row it was shown, the way a guardrail
|
||||
that scans per message does, and optionally the rewritten rows themselves."""
|
||||
|
||||
def post(url: str, json: dict, headers: dict) -> MagicMock:
|
||||
rows = json["structured_messages"]
|
||||
answer: dict = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": [_redacted(row["content"]) if isinstance(row.get("content"), str) else "" for row in rows],
|
||||
}
|
||||
if structured_messages_in_answer:
|
||||
answer["structured_messages"] = [{**row, "content": _redacted(row.get("content"))} for row in rows]
|
||||
response = MagicMock()
|
||||
response.json.return_value = answer
|
||||
response.raise_for_status = MagicMock()
|
||||
return response
|
||||
|
||||
return post
|
||||
|
||||
|
||||
def _per_message_redactor() -> GenericGuardrailAPI:
|
||||
return GenericGuardrailAPI(
|
||||
api_base="https://guardrail.test",
|
||||
guardrail_name="per-message-redactor",
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
|
||||
def _tool_replay_request() -> dict:
|
||||
return {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Never repeat the SSN " + SSN + " back.",
|
||||
"input": [
|
||||
{"role": "user", "content": "Look up " + SSN + " for me."},
|
||||
{"type": "function_call", "call_id": "call_1", "name": "lookup_customer", "arguments": '{"id": "42"}'},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _string_input_request() -> dict:
|
||||
return {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Never repeat the SSN " + SSN + " back.",
|
||||
"input": "My SSN is " + SSN + ".",
|
||||
}
|
||||
|
||||
|
||||
class TestPerMessageRewriteWriteBack:
|
||||
"""A guardrail that rewrites per chat row hands the rows back as
|
||||
structured_messages, and the handler lands them on the instructions and the
|
||||
input items they came from; the same rewrite handed back as texts alone has
|
||||
no item to land on and is rejected by name instead of sent unrewritten."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_structured_rows_land_on_instructions_and_tool_output(self):
|
||||
guardrail = _per_message_redactor()
|
||||
data = _tool_replay_request()
|
||||
function_call_item = data["input"][1]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)):
|
||||
result = await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back."
|
||||
assert _texts(result["input"][0]) == ["Look up " + REDACTED_SSN + " for me."]
|
||||
assert result["input"][1] == function_call_item
|
||||
assert result["input"][2] == {
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_1",
|
||||
"output": '{"ssn": "' + REDACTED_SSN + '"}',
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_texts_only_per_message_answer_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
guardrail = _per_message_redactor()
|
||||
data = _tool_replay_request()
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)):
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-message-redactor"
|
||||
assert data["input"] == original["input"]
|
||||
assert data["instructions"] == original["instructions"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_structured_rows_land_on_instructions_and_string_input(self):
|
||||
guardrail = _per_message_redactor()
|
||||
data = _string_input_request()
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)):
|
||||
result = await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back."
|
||||
assert [_texts(item) for item in result["input"]] == [["My SSN is " + REDACTED_SSN + "."]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_texts_only_per_message_answer_over_a_string_input_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
guardrail = _per_message_redactor()
|
||||
data = _string_input_request()
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)):
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-message-redactor"
|
||||
assert data["input"] == original["input"]
|
||||
assert data["instructions"] == original["instructions"]
|
||||
|
||||
|
||||
class TestProvenancePatching:
|
||||
"""The O(n) provenance pass must keep patching rewritten rows in place for the
|
||||
shapes real agent loops produce, and fall back safely everywhere else."""
|
||||
|
|
|
|||
|
|
@ -17,9 +17,9 @@ from unittest.mock import patch
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402
|
||||
VertexAIBatchTransformation,
|
||||
vertex_prompt_tokens_details,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import ( # noqa: E402
|
||||
VertexAIError,
|
||||
|
|
@ -41,6 +41,22 @@ ENDPOINT_INPUT_FILE = (
|
|||
)
|
||||
|
||||
|
||||
def test_vertex_prompt_tokens_details_rejects_malformed_details():
|
||||
assert vertex_prompt_tokens_details({"promptTokensDetails": [1]}) is None
|
||||
assert vertex_prompt_tokens_details({"promptTokensDetails": [{"modality": "AUDIO"}]}) is None
|
||||
assert (
|
||||
vertex_prompt_tokens_details(
|
||||
{
|
||||
"promptTokensDetails": [
|
||||
{"modality": "AUDIO", "tokenCount": 1},
|
||||
"malformed",
|
||||
]
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# transform_openai_batch_request_to_vertex_ai_batch_request
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ Covers:
|
|||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation import (
|
||||
_build_part_for_input,
|
||||
|
|
@ -22,11 +23,19 @@ from litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation
|
|||
from litellm.types.llms.vertex_ai import VertexAIBatchEmbeddingsResponseObject
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
||||
IMAGE_DATA_URI = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAIAQMAAAD+wSzIAAAABlBMVEX///+/v7+jQ3Y5AAAADklEQVQI12P4AIX8EAgALgAD/aNpbtEAAAAASUVORK5CYII"
|
||||
GCS_URL = "gs://my-bucket/image.png"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _local_model_cost_map(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
|
||||
litellm.get_model_info.cache_clear()
|
||||
yield
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
class TestIsMultimodalInput:
|
||||
def test_text_only_string(self):
|
||||
assert _is_multimodal_input("hello world") is False
|
||||
|
|
@ -324,7 +333,7 @@ class TestProcessEmbedContentResponseUsage:
|
|||
)
|
||||
assert result.usage.prompt_tokens == 258
|
||||
assert result.usage.total_tokens == 258
|
||||
assert result.usage.prompt_tokens_details.image_count == 1
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 258
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
|
|
@ -358,7 +367,7 @@ class TestProcessEmbedContentResponseUsage:
|
|||
)
|
||||
assert prompt_cost > 0
|
||||
|
||||
def test_video_modality_derives_seconds_and_text_floor(self):
|
||||
def test_video_modality_preserves_token_count(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
|
|
@ -374,10 +383,8 @@ class TestProcessEmbedContentResponseUsage:
|
|||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens == 516
|
||||
assert result.usage.prompt_tokens_details.video_length_seconds == pytest.approx(
|
||||
2.0
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 1
|
||||
assert result.usage.prompt_tokens_details.video_tokens == 516
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
|
||||
def test_missing_usage_metadata_does_not_estimate_from_base64(self):
|
||||
response_json = {"embedding": {"values": [0.1, 0.2]}}
|
||||
|
|
@ -400,8 +407,7 @@ class TestProcessEmbedContentResponseUsage:
|
|||
)
|
||||
assert result.usage.prompt_tokens > 0
|
||||
|
||||
def test_file_reference_image_billed_per_image_not_text(self):
|
||||
"""files/... image refs must bill per-image, not at the text token rate."""
|
||||
def test_file_reference_image_billed_per_image_token_rate(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1, 0.2, 0.3]},
|
||||
"usageMetadata": {
|
||||
|
|
@ -422,7 +428,7 @@ class TestProcessEmbedContentResponseUsage:
|
|||
}
|
||||
},
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_count == 1
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 258
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
|
|
@ -430,10 +436,10 @@ class TestProcessEmbedContentResponseUsage:
|
|||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(0.00012)
|
||||
assert prompt_cost == pytest.approx(258 * 4.5e-7)
|
||||
|
||||
def test_file_reference_non_image_not_counted_as_image(self):
|
||||
"""A files/... ref resolving to a non-image mime must not be image-counted."""
|
||||
"""A files/... ref resolving to a non-image mime keeps audio token billing."""
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1, 0.2]},
|
||||
"usageMetadata": {
|
||||
|
|
@ -454,21 +460,18 @@ class TestProcessEmbedContentResponseUsage:
|
|||
}
|
||||
},
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_count == 0
|
||||
assert result.usage.prompt_tokens_details.audio_tokens == 64
|
||||
assert result.usage.prompt_tokens_details.audio_length_seconds == pytest.approx(
|
||||
2.0
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(2.0 * 0.00016)
|
||||
assert prompt_cost == pytest.approx(64 * 6.5e-6)
|
||||
|
||||
def test_video_plus_audio_does_not_double_bill_text(self):
|
||||
"""Video+audio responses must not get video tokens reassigned to text."""
|
||||
"""Video and audio responses are billed from their respective token counts."""
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
|
|
@ -486,18 +489,145 @@ class TestProcessEmbedContentResponseUsage:
|
|||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 1
|
||||
assert result.usage.prompt_tokens_details.video_length_seconds == pytest.approx(
|
||||
2.0
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.audio_length_seconds == pytest.approx(
|
||||
2.0
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
assert result.usage.prompt_tokens_details.video_tokens == 516
|
||||
assert result.usage.prompt_tokens_details.audio_tokens == 64
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
# 1 floor text token at 2e-7 + 2s of video at 7.9e-4 + 2s of audio at 1.6e-4
|
||||
assert prompt_cost == pytest.approx(1 * 2e-7 + 2 * 0.00079 + 2 * 0.00016)
|
||||
assert prompt_cost == pytest.approx(516 * 1.2e-5 + 64 * 6.5e-6)
|
||||
|
||||
def test_preview_alias_bills_audio_per_token(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 64,
|
||||
"totalTokenCount": 64,
|
||||
"promptTokensDetails": [{"modality": "AUDIO", "tokenCount": 64}],
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input="audio",
|
||||
model_response=EmbeddingResponse(),
|
||||
model="gemini-embedding-2-preview",
|
||||
response_json=response_json,
|
||||
)
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model="gemini-embedding-2-preview",
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(64 * 6.5e-6)
|
||||
|
||||
def test_image_without_modality_details_uses_image_rate(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 258,
|
||||
"totalTokenCount": 258,
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input=IMAGE_DATA_URI,
|
||||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 258
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(258 * 4.5e-7)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"input_value,resolved_files,expected_image_tokens",
|
||||
[
|
||||
(GCS_URL, {}, 258),
|
||||
("gs://my-bucket/clip.mp4", {}, 0),
|
||||
("gs://my-bucket/unknown.bin", {}, 0),
|
||||
("files/image-123", {"files/image-123": {"mime_type": "image/jpeg"}}, 258),
|
||||
("files/missing", {}, 0),
|
||||
("data:application/octet-stream;base64,abc", {}, 0),
|
||||
([[IMAGE_DATA_URI]], {}, 258),
|
||||
([], {}, 0),
|
||||
],
|
||||
)
|
||||
def test_missing_modality_details_classifies_image_inputs(self, input_value, resolved_files, expected_image_tokens):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 258,
|
||||
"totalTokenCount": 258,
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input=input_value,
|
||||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
resolved_files=resolved_files,
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_tokens == expected_image_tokens
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
expected_rate = 4.5e-7 if expected_image_tokens else 2e-7
|
||||
assert prompt_cost == pytest.approx(258 * expected_rate)
|
||||
|
||||
def test_mixed_text_and_image_without_modality_details_not_billed_as_image(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 270,
|
||||
"totalTokenCount": 270,
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input=["a short caption", IMAGE_DATA_URI],
|
||||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(270 * 2e-7)
|
||||
|
||||
def test_text_without_modality_details_uses_text_rate(self):
|
||||
response_json = {
|
||||
"embedding": {"values": [0.1]},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 12,
|
||||
"totalTokenCount": 12,
|
||||
},
|
||||
}
|
||||
result = process_embed_content_response(
|
||||
input="a short caption",
|
||||
model_response=EmbeddingResponse(),
|
||||
model=self.MODEL,
|
||||
response_json=response_json,
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.text_tokens == 0
|
||||
assert result.usage.prompt_tokens_details.image_tokens == 0
|
||||
|
||||
prompt_cost, _ = generic_cost_per_token(
|
||||
model=self.MODEL,
|
||||
usage=result.usage,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert prompt_cost == pytest.approx(12 * 2e-7)
|
||||
|
|
|
|||
|
|
@ -1820,7 +1820,7 @@ async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(
|
|||
Skipping the write-back would hand the model the unredacted text, so a
|
||||
guardrail could be bypassed by adding ``instructions`` or a tool call.
|
||||
"""
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
data: dict[str, object] = {"model": "gpt-4o", "input": responses_input}
|
||||
if instructions is not None:
|
||||
|
|
|
|||
|
|
@ -582,6 +582,145 @@ class TestGuardrailActions:
|
|||
assert result_images is None
|
||||
|
||||
|
||||
class TestStructuredMessagesInResponse:
|
||||
"""A guardrail server that rewrites per chat row answers with the rewritten
|
||||
rows as structured_messages, which the endpoint handlers write back by row."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returned_rows_are_handed_back_as_structured_messages(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
rewritten_rows = [
|
||||
{"role": "system", "content": "Never repeat an SSN."},
|
||||
{"role": "user", "content": "Look up [REDACTED] for me."},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'},
|
||||
]
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["Never repeat an SSN.", "Look up [REDACTED] for me.", '{"ssn": "[REDACTED]"}'],
|
||||
"structured_messages": rewritten_rows,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."]},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert guardrailed_inputs["structured_messages"] == rewritten_rows
|
||||
assert guardrailed_inputs["texts"] == mock_response.json.return_value["texts"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rows_echoed_back_as_shown_keep_their_original_keys(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
tool_call_row = {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}, "index": 0}
|
||||
],
|
||||
}
|
||||
original_rows = [
|
||||
{"role": "user", "content": "Look up 123-45-6789 for me.", "name": "pat"},
|
||||
tool_call_row,
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'},
|
||||
]
|
||||
|
||||
def echo_with_tool_output_redacted(url, json, headers):
|
||||
shown_rows = json["structured_messages"]
|
||||
assert "index" not in shown_rows[1]["tool_calls"][0]
|
||||
assert "name" not in shown_rows[0]
|
||||
answer = MagicMock()
|
||||
answer.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["Look up 123-45-6789 for me."],
|
||||
"structured_messages": [
|
||||
shown_rows[0],
|
||||
shown_rows[1],
|
||||
{**shown_rows[2], "content": '{"ssn": "[REDACTED]"}'},
|
||||
],
|
||||
}
|
||||
answer.raise_for_status = MagicMock()
|
||||
return answer
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_with_tool_output_redacted):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."], "structured_messages": original_rows},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
returned_rows = guardrailed_inputs["structured_messages"]
|
||||
assert returned_rows[0] is original_rows[0]
|
||||
assert returned_rows[1] is tool_call_row
|
||||
assert returned_rows[2] == {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rows_all_echoed_back_as_shown_leave_the_rewrite_to_texts(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
"""A server written against the texts contract that echoes the request rows back
|
||||
untouched while rewriting texts still gets its texts rewrite applied."""
|
||||
original_rows = [
|
||||
{"role": "system", "content": "Never repeat an SSN."},
|
||||
{"role": "user", "content": "Look up 123-45-6789 for me."},
|
||||
]
|
||||
|
||||
def echo_rows_and_rewrite_texts(url, json, headers):
|
||||
answer = MagicMock()
|
||||
answer.json.return_value = {
|
||||
"action": "NONE",
|
||||
"texts": [text.replace("123-45-6789", "[REDACTED]") for text in json["texts"]],
|
||||
"structured_messages": json["structured_messages"],
|
||||
}
|
||||
answer.raise_for_status = MagicMock()
|
||||
return answer
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_rows_and_rewrite_texts):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={
|
||||
"texts": ["Never repeat an SSN.", "Look up 123-45-6789 for me."],
|
||||
"structured_messages": original_rows,
|
||||
},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert "structured_messages" not in guardrailed_inputs
|
||||
assert guardrailed_inputs["texts"] == ["Never repeat an SSN.", "Look up [REDACTED] for me."]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"structured_messages",
|
||||
[[], [{"content": "a row with no role"}], "not a list"],
|
||||
ids=["empty", "no_role", "not_a_list"],
|
||||
)
|
||||
async def test_rows_that_are_not_chat_messages_are_ignored(
|
||||
self, generic_guardrail, mock_request_data_input, structured_messages
|
||||
):
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["[REDACTED]"],
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."]},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert "structured_messages" not in guardrailed_inputs
|
||||
assert guardrailed_inputs["texts"] == ["[REDACTED]"]
|
||||
|
||||
|
||||
class TestImageSupport:
|
||||
"""Test image handling in guardrail requests"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1797,12 +1797,8 @@ PARTS_MESSAGES = [
|
|||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Earlier turn.", "cache_control": {"type": "ephemeral"}},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Second block. " + "B" * 5000,
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h"},
|
||||
},
|
||||
{"type": "text", "text": "Earlier turn."},
|
||||
{"type": "text", "text": "Second block. " + "B" * 5000},
|
||||
],
|
||||
},
|
||||
{
|
||||
|
|
@ -1891,14 +1887,9 @@ async def test_apply_guardrail_restores_rewritten_all_text_row(
|
|||
|
||||
messages = result["structured_messages"]
|
||||
history_content = messages[1]["content"]
|
||||
# Rewritten all-text row collapses to one part carrying the LAST declared
|
||||
# breakpoint: an Anthropic breakpoint caches the prefix ending at its
|
||||
# part, so after the merge the last one (and its TTL) still describes the
|
||||
# row.
|
||||
assert isinstance(history_content, list)
|
||||
assert len(history_content) == 1
|
||||
assert history_content[0]["text"] == "compressed history. Retrieve more: hash=b573993006976af767214fac"
|
||||
assert history_content[0]["cache_control"] == {"type": "ephemeral", "ttl": "1h"}
|
||||
# Mixed row passes through byte-identical.
|
||||
assert messages[2]["content"] == PARTS_MESSAGES[2]["content"]
|
||||
# The service-declared hash still drives retrieve-tool injection on a restored row.
|
||||
|
|
@ -2523,6 +2514,35 @@ async def test_history_is_still_compressed(guardrail: HeadroomGuardrail):
|
|||
assert messages[3] == compressed_history[1]
|
||||
|
||||
|
||||
CACHE_MARKED_HISTORY_MESSAGES = [
|
||||
{"role": "system", "content": "You are Claude Code. " + "S" * 5000},
|
||||
{"role": "user", "content": "old question " + "Q" * 5000},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Reading the file now.",
|
||||
"tool_calls": [{"id": "old_1", "type": "function", "function": {"name": "Read", "arguments": "{}"}}],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "old_1",
|
||||
"content": [{"type": "text", "text": "large cached file body " + "F" * 5000}],
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
},
|
||||
{"role": "assistant", "content": "Summarized the file for you."},
|
||||
{"role": "user", "content": "live instruction"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mid_history_cache_control_row_is_never_sent_for_compression(guardrail: HeadroomGuardrail):
|
||||
wire, result = await _wire_and_result(guardrail, CACHE_MARKED_HISTORY_MESSAGES)
|
||||
|
||||
cached_row = CACHE_MARKED_HISTORY_MESSAGES[3]
|
||||
assert cached_row not in wire
|
||||
assert not any(row.get("tool_call_id") == "old_1" for row in wire)
|
||||
assert result["structured_messages"][3] == cached_row
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# #38558: a client that runs its own tool loop (e.g. Claude Code via the MCP
|
||||
# gateway) executes headroom_retrieve and echoes the recovered original content
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import base64
|
||||
from collections.abc import Mapping, Sequence
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -12,6 +13,7 @@ from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security im
|
|||
PromptSecurityGuardrailMissingSecrets,
|
||||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
|
||||
|
|
@ -174,6 +176,123 @@ async def test_apply_guardrail_modify_request(monkeypatch: pytest.MonkeyPatch):
|
|||
assert result["texts"] == ["User prompt with PII: SSN [REDACTED]"]
|
||||
|
||||
|
||||
def _modify_response(modified_messages: Sequence[Mapping[str, object]]) -> Response:
|
||||
mock_response = Response(
|
||||
json={"result": {"prompt": {"action": "modify", "modified_messages": modified_messages}}},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/protect"),
|
||||
)
|
||||
mock_response.raise_for_status = lambda: None
|
||||
return mock_response
|
||||
|
||||
|
||||
def _tool_replay_messages() -> list[AllMessageValues]:
|
||||
return [
|
||||
{"role": "system", "content": "Never echo an SSN like 123-45-6789."},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Look up 123-45-6789"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/id-card.png"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'},
|
||||
{"role": "user", "content": "Summarize what you found."},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_returns_structured_messages_with_tool_rows_kept(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A per-message modify verdict comes back as structured_messages so the
|
||||
endpoint handler can write it back by message, with the rows Prompt Security
|
||||
never saw (tool results) and the non-text parts (images) left in place."""
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True)
|
||||
messages = _tool_replay_messages()
|
||||
inputs = {"texts": ["Look up 123-45-6789", "Summarize what you found."], "structured_messages": messages}
|
||||
modified_messages = [
|
||||
{"role": "system", "content": "Never echo an SSN like [REDACTED]."},
|
||||
{"role": "user", "content": [{"type": "text", "text": "Look up [REDACTED]"}]},
|
||||
{"role": "assistant", "content": None},
|
||||
{"role": "user", "content": "Summarize what you found."},
|
||||
]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == [
|
||||
{"role": "system", "content": "Never echo an SSN like [REDACTED]."},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Look up [REDACTED]"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/id-card.png"}},
|
||||
],
|
||||
},
|
||||
messages[2],
|
||||
messages[3],
|
||||
{"role": "user", "content": "Summarize what you found."},
|
||||
]
|
||||
assert result["structured_messages"] is not messages
|
||||
assert result["texts"] == [
|
||||
"Never echo an SSN like [REDACTED].",
|
||||
"Look up [REDACTED]",
|
||||
"Summarize what you found.",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_with_unexpected_message_count_keeps_texts_only(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True)
|
||||
messages = _tool_replay_messages()
|
||||
inputs = {"texts": ["Look up 123-45-6789", "Summarize what you found."], "structured_messages": messages}
|
||||
modified_messages = [{"role": "user", "content": "Look up [REDACTED]"}]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] is messages
|
||||
assert result["texts"] == ["Look up [REDACTED]"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_keeps_empty_text_parts_as_slots(monkeypatch: pytest.MonkeyPatch):
|
||||
"""The chat handler counts an empty text part as a slot, so a modify verdict
|
||||
that echoes the empty part still lines up with the row and its texts."""
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True)
|
||||
messages: list[AllMessageValues] = [
|
||||
{"role": "user", "content": [{"type": "text", "text": "Look up 123-45-6789"}, {"type": "text", "text": ""}]}
|
||||
]
|
||||
inputs = {"texts": ["Look up 123-45-6789", ""], "structured_messages": messages}
|
||||
modified_messages = [
|
||||
{"role": "user", "content": [{"type": "text", "text": "Look up [REDACTED]"}, {"type": "text", "text": ""}]}
|
||||
]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == modified_messages
|
||||
assert result["texts"] == ["Look up [REDACTED]", ""]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_allow_request(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Test that apply_guardrail allows safe prompts"""
|
||||
|
|
|
|||
|
|
@ -3909,6 +3909,57 @@ def _batch_cache_usage() -> Usage:
|
|||
)
|
||||
|
||||
|
||||
def test_batch_cost_calculator_prices_multimodal_tokens_at_modality_rates():
|
||||
from litellm.cost_calculator import batch_cost_calculator
|
||||
|
||||
model_info: ModelInfo = {
|
||||
"input_cost_per_token_batches": 1e-7,
|
||||
"input_cost_per_audio_token_batches": 3.25e-6,
|
||||
"input_cost_per_image_token_batches": 2.25e-7,
|
||||
"input_cost_per_video_token_batches": 6e-6,
|
||||
}
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=0,
|
||||
total_tokens=100,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
audio_tokens=64,
|
||||
image_tokens=10,
|
||||
video_tokens=6,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, _ = batch_cost_calculator(
|
||||
usage=usage,
|
||||
model="gemini-embedding-2",
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_info=model_info,
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(20 * 1e-7 + 64 * 3.25e-6 + 10 * 2.25e-7 + 6 * 6e-6)
|
||||
|
||||
|
||||
def test_batch_cost_calculator_falls_back_to_text_batch_rate_for_modalities():
|
||||
from litellm.cost_calculator import batch_cost_calculator
|
||||
|
||||
model_info: ModelInfo = {"input_cost_per_token_batches": 1e-7}
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=0,
|
||||
total_tokens=100,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=64),
|
||||
)
|
||||
|
||||
prompt_cost, _ = batch_cost_calculator(
|
||||
usage=usage,
|
||||
model="gemini-embedding-2",
|
||||
custom_llm_provider="vertex_ai",
|
||||
model_info=model_info,
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(100 * 1e-7)
|
||||
|
||||
|
||||
def test_batch_cost_calculator_prices_cache_creation_tokens_at_cache_write_rate():
|
||||
"""
|
||||
LIT-4008 regression: anthropic batch usage is dominated by cache tokens.
|
||||
|
|
|
|||
|
|
@ -892,7 +892,10 @@ def validate_model_cost_values(model_data, exceptions=None):
|
|||
"input_cost_per_video_per_second_above_8s_interval",
|
||||
"input_cost_per_video_per_second_above_15s_interval",
|
||||
"input_cost_per_video_per_second_above_128k_tokens",
|
||||
"input_cost_per_audio_token_batches",
|
||||
"input_cost_per_image_token_batches",
|
||||
"input_cost_per_token_batches",
|
||||
"input_cost_per_video_token_batches",
|
||||
"output_cost_per_token_batches",
|
||||
"input_cost_per_token_cache_hit",
|
||||
"cache_creation_input_token_cost",
|
||||
|
|
@ -1041,7 +1044,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"input_cost_per_second": {"type": "number"},
|
||||
"input_cost_per_token": {"type": "number"},
|
||||
"input_cost_per_token_above_128k_tokens": {"type": "number"},
|
||||
"input_cost_per_audio_token_batches": {"type": "number"},
|
||||
"input_cost_per_image_token_batches": {"type": "number"},
|
||||
"input_cost_per_token_batches": {"type": "number"},
|
||||
"input_cost_per_video_token_batches": {"type": "number"},
|
||||
"input_cost_per_token_cache_hit": {"type": "number"},
|
||||
"input_cost_per_video_per_second": {"type": "number"},
|
||||
"input_cost_per_video_per_second_above_8s_interval": {"type": "number"},
|
||||
|
|
@ -2946,7 +2952,7 @@ def test_model_info_for_openrouter_kimi_k2_5():
|
|||
|
||||
|
||||
def test_gemini_embedding_2_ga_in_cost_map():
|
||||
"""GA and Vertex preview gemini-embedding-2 entries align with multimodal unit pricing."""
|
||||
"""GA and Vertex preview gemini-embedding-2 entries align with multimodal token pricing."""
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
|
|
@ -2968,9 +2974,15 @@ def test_gemini_embedding_2_ga_in_cost_map():
|
|||
assert info.get("mode") == "embedding"
|
||||
assert info.get("supports_multimodal") is True
|
||||
assert info.get("input_cost_per_token") == 2e-07
|
||||
assert info.get("input_cost_per_image") == 0.00012
|
||||
assert info.get("input_cost_per_audio_per_second") == 0.00016
|
||||
assert info.get("input_cost_per_video_per_second") == 0.00079
|
||||
assert info.get("input_cost_per_audio_token") == 6.5e-06
|
||||
assert info.get("input_cost_per_image_token") == 4.5e-07
|
||||
assert info.get("input_cost_per_video_token") == 1.2e-05
|
||||
assert info.get("input_cost_per_audio_token_batches") == 3.25e-06
|
||||
assert info.get("input_cost_per_image_token_batches") == 2.25e-07
|
||||
assert info.get("input_cost_per_video_token_batches") == 6e-06
|
||||
assert "input_cost_per_image" not in info
|
||||
assert "input_cost_per_audio_per_second" not in info
|
||||
assert "input_cost_per_video_per_second" not in info
|
||||
if provider in ("vertex_ai-embedding-models", "vertex_ai"):
|
||||
assert (
|
||||
info.get("uses_embed_content") is True
|
||||
|
|
|
|||
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -29857,6 +29857,8 @@ export interface components {
|
|||
input_cost_per_audio_per_second_above_128k_tokens?: number | null;
|
||||
/** Input Cost Per Audio Token */
|
||||
input_cost_per_audio_token?: number | null;
|
||||
/** Input Cost Per Audio Token Batches */
|
||||
input_cost_per_audio_token_batches?: number | null;
|
||||
/** Input Cost Per Character */
|
||||
input_cost_per_character?: number | null;
|
||||
/** Input Cost Per Character Above 128K Tokens */
|
||||
|
|
@ -29867,6 +29869,8 @@ export interface components {
|
|||
input_cost_per_image_above_128k_tokens?: number | null;
|
||||
/** Input Cost Per Image Token */
|
||||
input_cost_per_image_token?: number | null;
|
||||
/** Input Cost Per Image Token Batches */
|
||||
input_cost_per_image_token_batches?: number | null;
|
||||
/** Input Cost Per Pixel */
|
||||
input_cost_per_pixel?: number | null;
|
||||
/** Input Cost Per Query */
|
||||
|
|
@ -29909,6 +29913,8 @@ export interface components {
|
|||
input_cost_per_video_per_second_above_8s_interval?: number | null;
|
||||
/** Input Cost Per Video Token */
|
||||
input_cost_per_video_token?: number | null;
|
||||
/** Input Cost Per Video Token Batches */
|
||||
input_cost_per_video_token_batches?: number | null;
|
||||
/** Itpm */
|
||||
itpm?: number | null;
|
||||
/** Keepalive Seconds */
|
||||
|
|
@ -40077,6 +40083,8 @@ export interface components {
|
|||
input_cost_per_audio_per_second_above_128k_tokens?: number | null;
|
||||
/** Input Cost Per Audio Token */
|
||||
input_cost_per_audio_token?: number | null;
|
||||
/** Input Cost Per Audio Token Batches */
|
||||
input_cost_per_audio_token_batches?: number | null;
|
||||
/** Input Cost Per Character */
|
||||
input_cost_per_character?: number | null;
|
||||
/** Input Cost Per Character Above 128K Tokens */
|
||||
|
|
@ -40087,6 +40095,8 @@ export interface components {
|
|||
input_cost_per_image_above_128k_tokens?: number | null;
|
||||
/** Input Cost Per Image Token */
|
||||
input_cost_per_image_token?: number | null;
|
||||
/** Input Cost Per Image Token Batches */
|
||||
input_cost_per_image_token_batches?: number | null;
|
||||
/** Input Cost Per Pixel */
|
||||
input_cost_per_pixel?: number | null;
|
||||
/** Input Cost Per Query */
|
||||
|
|
@ -40129,6 +40139,8 @@ export interface components {
|
|||
input_cost_per_video_per_second_above_8s_interval?: number | null;
|
||||
/** Input Cost Per Video Token */
|
||||
input_cost_per_video_token?: number | null;
|
||||
/** Input Cost Per Video Token Batches */
|
||||
input_cost_per_video_token_batches?: number | null;
|
||||
/** Itpm */
|
||||
itpm?: number | null;
|
||||
/** Keepalive Seconds */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue