mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
test: deflake JWT tamper assertions and fuzzy picker widget driver
Tamper tests rewrote the last two base64url characters of the signature,
which on roughly 1 in 250 RS256 tokens (1 in 1000 HS256) only touched
padding bits, so the decoded signature was unchanged and still verified.
Corrupt the decoded signature bytes instead.
The fuzzy picker driver sent keys after fixed sleeps, so a slow worker
could receive the filter text before the widget had highlighted the match.
Wait on the widget's highlighted choice instead.
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
(cherry picked from commit 25c5f0d993)
This commit is contained in:
parent
e5f93c008a
commit
3651a47ee9
3 changed files with 53 additions and 24 deletions
|
|
@ -5,6 +5,7 @@ from datetime import datetime, timedelta, timezone
|
|||
import pytest
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from jwt.utils import base64url_decode, base64url_encode
|
||||
from pydantic import SecretStr
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import (
|
||||
|
|
@ -52,6 +53,12 @@ def _refresh_token() -> str:
|
|||
return minted.token.get_secret_value()
|
||||
|
||||
|
||||
def _corrupt_signature(token: str) -> str:
|
||||
unsigned, signature = token.rsplit(".", 1)
|
||||
raw = base64url_decode(signature)
|
||||
return f"{unsigned}.{base64url_encode(bytes((raw[0] ^ 0x01,)) + raw[1:]).decode()}"
|
||||
|
||||
|
||||
def test_kdf_is_deterministic_and_key_length_is_256_bit():
|
||||
again = session_keys_from_master_key(MASTER_KEY)
|
||||
assert again.signing_key.get_secret_value() == KEYS.signing_key.get_secret_value()
|
||||
|
|
@ -109,8 +116,7 @@ def test_resolve_fails_expired_token_closed_and_flags_expiry():
|
|||
|
||||
def test_resolve_fails_tampered_token_closed_without_expiry_flag():
|
||||
token = _access_token()
|
||||
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
|
||||
result = resolve_session_bearer(f"Bearer {tampered}", KEYS, NOW)
|
||||
result = resolve_session_bearer(f"Bearer {_corrupt_signature(token)}", KEYS, NOW)
|
||||
assert isinstance(result, SessionBearerInvalid)
|
||||
assert result.expired is False
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import jwt
|
|||
import pytest
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from jwt.utils import base64url_decode, base64url_encode
|
||||
from pydantic import SecretStr, ValidationError
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
|
||||
|
|
@ -66,6 +67,12 @@ def _mint_refresh() -> str:
|
|||
return minted.token.get_secret_value()
|
||||
|
||||
|
||||
def _corrupt_signature(token: str) -> str:
|
||||
unsigned, signature = token.rsplit(".", 1)
|
||||
raw = base64url_decode(signature)
|
||||
return f"{unsigned}.{base64url_encode(bytes((raw[0] ^ 0x01,)) + raw[1:]).decode()}"
|
||||
|
||||
|
||||
def _sign_claims(payload: dict, prefix: str = SESSION_TOKEN_PREFIX, keys: SessionKeys = KEYS) -> str:
|
||||
return prefix + jwt.encode(payload, keys.signing_key.get_secret_value(), algorithm="HS256")
|
||||
|
||||
|
|
@ -138,8 +145,7 @@ def test_still_valid_one_second_before_expiry():
|
|||
|
||||
def test_tampered_signature_is_bad_signature():
|
||||
token = _mint_access()
|
||||
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
|
||||
assert isinstance(open_session_token(tampered, KEYS, NOW), SessionBadSignature)
|
||||
assert isinstance(open_session_token(_corrupt_signature(token), KEYS, NOW), SessionBadSignature)
|
||||
|
||||
|
||||
def test_key_rotation_invalidates_outstanding_tokens():
|
||||
|
|
@ -329,8 +335,7 @@ def test_rs256_tampered_signature_is_bad_signature():
|
|||
minted = mint_session_token(PRINCIPAL, RSA_KEYS, NOW)
|
||||
assert isinstance(minted, MintedSessionToken)
|
||||
token = minted.token.get_secret_value()
|
||||
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
|
||||
assert isinstance(open_session_token(tampered, RSA_KEYS, NOW), SessionBadSignature)
|
||||
assert isinstance(open_session_token(_corrupt_signature(token), RSA_KEYS, NOW), SessionBadSignature)
|
||||
|
||||
|
||||
def test_rs256_expired_token_is_expired():
|
||||
|
|
@ -413,8 +418,7 @@ def test_rotation_window_still_enforces_expiry_and_tamper_on_the_previous_key():
|
|||
)
|
||||
after = NOW + timedelta(seconds=SESSION_TTL_SECONDS + 1)
|
||||
assert isinstance(open_session_token(token, rotated, after), SessionExpired)
|
||||
tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb")
|
||||
assert isinstance(open_session_token(tampered, rotated, NOW), SessionBadSignature)
|
||||
assert isinstance(open_session_token(_corrupt_signature(token), rotated, NOW), SessionBadSignature)
|
||||
|
||||
|
||||
def test_weak_or_garbage_private_key_pem_rejected_at_construction():
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import asyncio
|
||||
from typing import Any, Dict, List, Tuple
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
from unittest.mock import patch
|
||||
|
||||
import click
|
||||
|
|
@ -7,7 +7,8 @@ import pytest
|
|||
import yaml
|
||||
from click.testing import CliRunner
|
||||
from InquirerPy.base.control import Choice
|
||||
from prompt_toolkit.application import create_app_session
|
||||
from InquirerPy.prompts.fuzzy import InquirerPyFuzzyControl
|
||||
from prompt_toolkit.application import AppSession, create_app_session
|
||||
from prompt_toolkit.input import create_pipe_input
|
||||
from prompt_toolkit.output import DummyOutput
|
||||
|
||||
|
|
@ -283,27 +284,45 @@ class TestRunConfigureWizardNotInteractive:
|
|||
assert not config_path.exists()
|
||||
|
||||
|
||||
def _highlighted_choice(session: AppSession) -> Optional[str]:
|
||||
if session.app is None:
|
||||
return None
|
||||
controls = [c for c in session.app.layout.find_all_controls() if isinstance(c, InquirerPyFuzzyControl)]
|
||||
if not controls or controls[0].choice_count == 0:
|
||||
return None
|
||||
return controls[0].selection["name"]
|
||||
|
||||
|
||||
async def _wait_until_highlighted(session: AppSession, name: str) -> None:
|
||||
async def _poll() -> None:
|
||||
while _highlighted_choice(session) != name:
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
await asyncio.wait_for(_poll(), timeout=5)
|
||||
|
||||
|
||||
def _drive_fuzzy_pick(
|
||||
models: Tuple[DiscoveredModel, ...],
|
||||
prompt_label: str,
|
||||
multiselect: bool,
|
||||
key_events: List[Tuple[str, float]],
|
||||
key_events: List[Tuple[str, Optional[str]]],
|
||||
) -> List[str]:
|
||||
"""Drives the real InquirerPy fuzzy prompt through prompt_toolkit's own test input/output,
|
||||
exercising the actual widget (filtering, tab-to-toggle, enter-to-confirm) rather than mocking
|
||||
it away. asyncio.to_thread propagates the create_app_session context into the worker thread
|
||||
running _fuzzy_pick's synchronous .execute() call."""
|
||||
running _fuzzy_pick's synchronous .execute() call. Each key event names the choice the widget
|
||||
must highlight before the next key is sent (None sends the next key immediately)."""
|
||||
|
||||
async def _run() -> List[str]:
|
||||
with create_pipe_input() as pipe_input:
|
||||
with create_app_session(input=pipe_input, output=DummyOutput()):
|
||||
with create_app_session(input=pipe_input, output=DummyOutput()) as session:
|
||||
task = asyncio.ensure_future(
|
||||
asyncio.to_thread(wizard_module._fuzzy_pick, models, prompt_label, multiselect)
|
||||
)
|
||||
await asyncio.sleep(0.05)
|
||||
for text, delay in key_events:
|
||||
for text, highlighted in key_events:
|
||||
pipe_input.send_text(text)
|
||||
await asyncio.sleep(delay)
|
||||
if highlighted is not None:
|
||||
await _wait_until_highlighted(session, highlighted)
|
||||
return await task
|
||||
|
||||
return asyncio.run(_run())
|
||||
|
|
@ -315,13 +334,13 @@ class TestFuzzyPickWidget:
|
|||
|
||||
def test_single_select_filters_and_returns_highlighted_match(self):
|
||||
result = _drive_fuzzy_pick(
|
||||
self._models(), "test", multiselect=False, key_events=[("model-13", 0.3), ("\r", 0.1)]
|
||||
self._models(), "test", multiselect=False, key_events=[("model-13", "model-13"), ("\r", None)]
|
||||
)
|
||||
assert result == ["model-13"]
|
||||
|
||||
def test_multiselect_requires_tab_to_toggle_before_enter(self):
|
||||
result = _drive_fuzzy_pick(
|
||||
self._models(), "test", multiselect=True, key_events=[("model-7", 0.3), ("\t", 0.1), ("\r", 0.1)]
|
||||
self._models(), "test", multiselect=True, key_events=[("model-7", "model-7"), ("\t", None), ("\r", None)]
|
||||
)
|
||||
assert result == ["model-7"]
|
||||
|
||||
|
|
@ -331,12 +350,12 @@ class TestFuzzyPickWidget:
|
|||
"test",
|
||||
multiselect=True,
|
||||
key_events=[
|
||||
("model-3", 0.3),
|
||||
("\t", 0.1),
|
||||
*[("\x7f", 0.02) for _ in range("model-3".__len__())],
|
||||
("model-15", 0.3),
|
||||
("\t", 0.1),
|
||||
("\r", 0.1),
|
||||
("model-3", "model-3"),
|
||||
("\t", None),
|
||||
("\x7f" * len("model-3"), None),
|
||||
("model-15", "model-15"),
|
||||
("\t", None),
|
||||
("\r", None),
|
||||
],
|
||||
)
|
||||
assert set(result) == {"model-3", "model-15"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue