mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
Merge pull request #40892 from BerriAI/litellm_jwt_management_callers
test: bind management E2E callers and isolate JWT actors
This commit is contained in:
commit
15bd8b0e4a
23 changed files with 1602 additions and 175 deletions
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-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[@]}"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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`:
|
||||
|
||||
|
|
|
|||
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:
|
||||
|
|
|
|||
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) ----------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
},
|
||||
});
|
||||
Loading…
Add table
Reference in a new issue