mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
test(integration): credential canary slots for MCP and pass-through credentials (#43308)
* test(integration): credential canary suite harness Adds tests/integration/security with canary generation and search, sweeps over the database, GET routes, client responses, sink doubles and Redis, an owned proxy rig, a sweep sensitivity self-test and the config deployment api_key slot. Registers the security group in run.py, the manifest and the CircleCI integration matrix. * test(integration): widen canary route sweep and harden the rig Enumerate lazily registered feature routers, call parameterized routes with placeholder ids, fail on routes that return no response, skip provider pass-through routes, add an explicit admin-only route allowance, let the sink double use a configurable token, inflate gzip members anywhere in a blob, sweep Redis before the route walk, and trap outbound connections from the owned proxy. * test(integration): descend into any decoded value that can still hold an encoded canary * test(integration): bound canary decoding by depth and decoded bytes * test(integration): scope log-table and spend-log reads to the scenario window * test(integration): sweep spend-log rows in the scenario date window * test(integration): keep spend-log date window summarized * test(integration): credential canary slots for MCP and pass-through credentials Adds slots F1 (MCP static auth), F2 (per-user MCP OAuth token), F2E (per-user MCP env var), F3 (x-mcp client auth header), H1 (pass-through credential header), H2 (vector store api_key) and H2S (search tool api_key) to the credential canary suite. The OAuth double gains an optional mint hook so a test can choose the issued access token. * test(integration): wait for MCP spend rows by call type * test(integration): resolve deployment ids, scope paginated log lists, key allowances by slot * test(integration): canary MCP and pass-through slots pass resolved ids * test(integration): expect 404 from the caller-scoped team membership route * test(integration): use the rig's own master key and expect 404 from submission lookups * test(integration): check the overridden rig key without assuming the default key is unknown
This commit is contained in:
parent
5e38a08741
commit
db9c307e1a
4 changed files with 588 additions and 4 deletions
|
|
@ -7,7 +7,7 @@ import json
|
|||
import secrets
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from collections.abc import Callable, Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Final
|
||||
|
|
@ -27,6 +27,7 @@ class AuthorizationServer:
|
|||
refresh_tokens: dict[str, dict[str, str]] = field(default_factory=dict)
|
||||
revoked: set[str] = field(default_factory=set)
|
||||
lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
mint: Callable[[str], str] | None = None
|
||||
|
||||
@property
|
||||
def issuer(self) -> str:
|
||||
|
|
@ -47,7 +48,7 @@ class AuthorizationServer:
|
|||
return token in self.access_tokens and token not in self.revoked
|
||||
|
||||
def issue(self, grant: str, client_id: str, subject: str, scope: str) -> dict[str, object]:
|
||||
access: Final = f"at-{grant}-{secrets.token_urlsafe(8)}"
|
||||
access: Final = self.mint(grant) if self.mint is not None else f"at-{grant}-{secrets.token_urlsafe(8)}"
|
||||
refresh: Final = f"rt-{secrets.token_urlsafe(8)}"
|
||||
with self.lock:
|
||||
self.access_tokens[access] = {"client_id": client_id, "subject": subject, "scope": scope, "grant": grant}
|
||||
|
|
@ -80,7 +81,10 @@ def _client_credentials(request: Request, form: dict[str, str]) -> tuple[str, st
|
|||
|
||||
|
||||
@contextmanager
|
||||
def oauth_server(*, scopes: tuple[str, ...] = ("tools.read", "tools.call")) -> Iterator[AuthorizationServer]:
|
||||
def oauth_server(
|
||||
*, scopes: tuple[str, ...] = ("tools.read", "tools.call"), mint: Callable[[str], str] | None = None
|
||||
) -> Iterator[AuthorizationServer]:
|
||||
"""``mint(grant)``, when given, chooses each issued access token instead of a random one."""
|
||||
holder: list[AuthorizationServer] = []
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
|
|
@ -194,5 +198,5 @@ def oauth_server(*, scopes: tuple[str, ...] = ("tools.read", "tools.call")) -> I
|
|||
return _json(404, {"error": "not_found", "path": path, "method": request.method})
|
||||
|
||||
with wire_server(respond) as wire:
|
||||
holder.append(AuthorizationServer(wire))
|
||||
holder.append(AuthorizationServer(wire, mint=mint))
|
||||
yield holder[0]
|
||||
|
|
|
|||
|
|
@ -81,6 +81,13 @@ SLOTS: Final = MappingProxyType(
|
|||
{
|
||||
MARKER: Slot(MARKER, "Sensitivity marker in message content; must appear where prompts are stored"),
|
||||
"B1": Slot("B1", "Deployment api_key declared in the proxy config.yaml model_list"),
|
||||
"F1": Slot("F1", "MCP server static auth_value registered through /v1/mcp/server"),
|
||||
"F2": Slot("F2", "Per-user MCP OAuth access token from the authorization-code flow"),
|
||||
"F2E": Slot("F2E", "Per-user MCP env var value stored through /v1/mcp/server/{server_id}/user-env-vars"),
|
||||
"F3": Slot("F3", "Client x-mcp-<server>-authorization request header"),
|
||||
"H1": Slot("H1", "Pass-through endpoint credential header resolved from os.environ"),
|
||||
"H2": Slot("H2", "Vector store api_key declared in the proxy config.yaml vector_store_registry"),
|
||||
"H2S": Slot("H2S", "Search tool api_key declared in the proxy config.yaml search_tools"),
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
363
tests/integration/security/test_mcp_slots.py
Normal file
363
tests/integration/security/test_mcp_slots.py
Normal file
|
|
@ -0,0 +1,363 @@
|
|||
"""Slots F1 to F3: MCP credentials reach only the MCP peer they belong to.
|
||||
|
||||
Each scenario registers a scripted MCP peer (``_support/mcp.py``) that records every request,
|
||||
wires one credential slot to it and calls a tool, either directly over the server's MCP
|
||||
endpoint or through ``/v1/chat/completions`` with the provider double asking for the tool. The
|
||||
``echo`` tool succeeds and the ``deny`` tool answers HTTP 401, so both the success and the
|
||||
upstream-rejection logging paths run.
|
||||
|
||||
- F1: static ``auth_value`` registered through ``/v1/mcp/server``.
|
||||
- F2: per-user OAuth access token, issued by the OAuth 2.1 double through the gateway's
|
||||
authorization-code flow with PKCE.
|
||||
- F2E: per-user env var value, stored through ``/v1/mcp/server/{server_id}/user-env-vars`` and
|
||||
substituted into the server's ``Authorization`` header.
|
||||
- F3: client ``x-mcp-<server>-authorization`` request header.
|
||||
|
||||
Positive control: the peer's ``tools/call`` request must carry ``Authorization: Bearer
|
||||
<canary>``, or the test fails before sweeping. Sensitivity control: the marker sent as the tool
|
||||
argument must be reported where stored prompts belong. Then no sweep may find the canary.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import secrets
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Scenario, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.mcp import JsonRpc, McpCaller, McpPeer, ScriptedTool, echo_tool, register_mcp, scripted_peer
|
||||
from integration._support.oauth_server import oauth_server
|
||||
from integration._support.wire import Reply, Request
|
||||
from integration.security._canary import MARKER, Canary, canary, find_canary
|
||||
from integration.security._sinks import CONFIG_MODEL, GENERIC_SINK, Caller, Rig, canary_rig, chat_upstream, settle
|
||||
from integration.security._sweeps import Hit, assert_marker_seen, assert_no_hits, record_route_sweep, sweep_all
|
||||
|
||||
Via = Literal["direct", "chat"]
|
||||
Outcome = Literal["success", "upstream_401"]
|
||||
TOOL: Final[Mapping[Outcome, str]] = {"success": "echo", "upstream_401": "deny"}
|
||||
USER_TOKEN: Final = "USER_TOKEN"
|
||||
CLIENT_REDIRECT: Final = "http://127.0.0.1:9/cb"
|
||||
SLACK: Final = timedelta(seconds=5)
|
||||
|
||||
|
||||
def _tool_call(request: Request) -> Reply:
|
||||
"""Provider double: asks for the first offered tool with the user text, then echoes the tool result."""
|
||||
body: Final = json.loads(request.body or b"{}")
|
||||
tools: Final = body.get("tools") or []
|
||||
messages: Final = body.get("messages") or []
|
||||
if not tools or any(message.get("role") == "tool" for message in messages):
|
||||
return chat_upstream(request)
|
||||
call: Final = {
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tools[0]["function"]["name"],
|
||||
"arguments": json.dumps({"text": str(messages[-1].get("content", ""))}),
|
||||
},
|
||||
}
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex}",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"finish_reason": "tool_calls",
|
||||
"message": {"role": "assistant", "content": None, "tool_calls": [call]},
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _deny(params: JsonRpc) -> Reply:
|
||||
return Reply(
|
||||
status=401,
|
||||
body=b'{"error":"invalid_token"}',
|
||||
headers={"www-authenticate": 'Bearer error="invalid_token"'},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
|
||||
"""One owned proxy per module: every F credential is registered at runtime with a fresh core."""
|
||||
with canary_rig(tmp_path_factory.mktemp("canary-mcp"), upstream=_tool_call) as value:
|
||||
yield value
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Wiring:
|
||||
server_id: str
|
||||
alias: str
|
||||
caller: Caller
|
||||
headers: Mapping[str, str] = field(default_factory=dict)
|
||||
responses: tuple[httpx.Response, ...] = ()
|
||||
|
||||
|
||||
def _caller(scenario: Scenario, server_id: str) -> Caller:
|
||||
grant: Final[JsonRpc] = {"mcp_servers": [server_id]}
|
||||
team: Final = scenario.team(object_permission=dict(grant))
|
||||
user: Final = scenario.user(user_role="internal_user")
|
||||
scenario.gateway.post("/team/member_add", {"team_id": team, "member": {"user_id": user, "role": "user"}})
|
||||
key: Final = scenario.key(team_id=team, user_id=user, models=[CONFIG_MODEL], object_permission=dict(grant))
|
||||
return Caller(team, user, key)
|
||||
|
||||
|
||||
def _pkce_challenge(verifier: str) -> str:
|
||||
return base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
|
||||
|
||||
|
||||
def _authorize_and_redeem(rig: Rig, alias: str, key: str) -> None:
|
||||
"""Run the gateway's authorization-code flow for the caller; the double mints the canary."""
|
||||
client: Final = rig.proxy.client
|
||||
base: Final = str(client.base_url).rstrip("/")
|
||||
registered: Final = client.post(f"/{alias}/register", json={"redirect_uris": [CLIENT_REDIRECT]})
|
||||
assert registered.status_code in (200, 201), registered.text
|
||||
client_id: Final = string_value(registered.json()["client_id"])
|
||||
verifier: Final = secrets.token_urlsafe(32)
|
||||
started: Final = client.get(
|
||||
f"/{alias}/authorize",
|
||||
params={
|
||||
"client_id": client_id,
|
||||
"redirect_uri": CLIENT_REDIRECT,
|
||||
"response_type": "code",
|
||||
"state": "canary-state",
|
||||
"code_challenge": _pkce_challenge(verifier),
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "tools.call",
|
||||
},
|
||||
headers={"x-litellm-api-key": key},
|
||||
)
|
||||
assert started.status_code in (302, 307), started.text
|
||||
consent: Final = httpx.get(started.headers["location"], follow_redirects=False, trust_env=False)
|
||||
assert consent.status_code == 302, consent.text
|
||||
returned: Final = client.get(
|
||||
consent.headers["location"].removeprefix(base), headers={"x-litellm-api-key": key}, cookies=started.cookies
|
||||
)
|
||||
assert returned.status_code == 302, returned.text
|
||||
code: Final = parse_qs(urlsplit(returned.headers["location"]).query)["code"][0]
|
||||
redeemed: Final = client.post(
|
||||
f"/{alias}/token",
|
||||
headers={"x-litellm-api-key": key},
|
||||
data={
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"code_verifier": verifier,
|
||||
"client_id": client_id,
|
||||
"redirect_uri": CLIENT_REDIRECT,
|
||||
},
|
||||
)
|
||||
assert redeemed.status_code == 200, redeemed.text
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _wired(slot: str, rig: Rig, scenario: Scenario, peer: McpPeer, credential: Canary) -> Iterator[Wiring]:
|
||||
"""Register the peer with ``credential`` in ``slot`` and return the caller that uses it."""
|
||||
alias: Final = "canary" + uuid.uuid4().hex[:8]
|
||||
if slot == "F1":
|
||||
server: Final = register_mcp(
|
||||
scenario, peer, alias, auth_type="bearer_token", credentials={"auth_value": credential.value}
|
||||
)
|
||||
yield Wiring(server, alias, _caller(scenario, server))
|
||||
elif slot == "F2":
|
||||
with oauth_server(mint=lambda grant: credential.value) as auth:
|
||||
server_f2: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="oauth2",
|
||||
oauth2_flow="authorization_code",
|
||||
issuer=auth.issuer,
|
||||
authorization_url=auth.issuer + "/authorize",
|
||||
token_url=auth.issuer + "/token",
|
||||
registration_url=auth.issuer + "/register",
|
||||
credentials={"client_id": "canary-client", "client_secret": "canary-client-secret"},
|
||||
)
|
||||
caller_f2: Final = _caller(scenario, server_f2)
|
||||
_authorize_and_redeem(rig, alias, caller_f2.key)
|
||||
yield Wiring(server_f2, alias, caller_f2)
|
||||
elif slot == "F2E":
|
||||
server_f2e: Final = register_mcp(
|
||||
scenario,
|
||||
peer,
|
||||
alias,
|
||||
auth_type="none",
|
||||
env_vars=[{"name": USER_TOKEN, "scope": "user", "description": "per-user token"}],
|
||||
static_headers={"Authorization": f"Bearer ${{{USER_TOKEN}}}"},
|
||||
)
|
||||
caller_f2e: Final = _caller(scenario, server_f2e)
|
||||
stored: Final = rig.proxy.request(
|
||||
"POST",
|
||||
f"/v1/mcp/server/{server_f2e}/user-env-vars",
|
||||
{"values": {USER_TOKEN: credential.value}},
|
||||
key=caller_f2e.key,
|
||||
)
|
||||
assert stored.status_code == 200, stored.text
|
||||
yield Wiring(server_f2e, alias, caller_f2e, responses=(stored,))
|
||||
else:
|
||||
assert slot == "F3", slot
|
||||
server_f3: Final = register_mcp(scenario, peer, alias)
|
||||
yield Wiring(
|
||||
server_f3,
|
||||
alias,
|
||||
_caller(scenario, server_f3),
|
||||
headers={f"x-mcp-{alias}-authorization": f"Bearer {credential.value}"},
|
||||
)
|
||||
|
||||
|
||||
def _send(rig: Rig, wiring: Wiring, via: Via, tool: str, text: str) -> httpx.Response:
|
||||
if via == "direct":
|
||||
return McpCaller(rig.proxy, wiring.caller.key, "server_mcp", wiring.alias, wiring.headers).rpc(
|
||||
"tools/call", {"name": f"{wiring.alias}-{tool}", "arguments": {"text": text}}
|
||||
)
|
||||
return rig.proxy.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": CONFIG_MODEL,
|
||||
"messages": [{"role": "user", "content": text}],
|
||||
"tools": [
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": f"litellm_proxy/mcp/{wiring.alias}",
|
||||
"server_label": "litellm",
|
||||
"require_approval": "never",
|
||||
"allowed_tools": [f"{wiring.alias}-{tool}"],
|
||||
}
|
||||
],
|
||||
},
|
||||
key=wiring.caller.key,
|
||||
headers=wiring.headers,
|
||||
)
|
||||
|
||||
|
||||
def _answer(response: httpx.Response, via: Via) -> str:
|
||||
"""The text the caller got back: the tool result (direct) or the assistant message (chat)."""
|
||||
assert response.status_code == 200, response.text
|
||||
if via == "chat":
|
||||
return string_value(object_value(response.json()["choices"][0]["message"])["content"])
|
||||
data: Final = next(
|
||||
line.removeprefix("data:").strip() for line in response.text.splitlines() if line.startswith("data:")
|
||||
)
|
||||
result: Final = object_value(json.loads(data)["result"])
|
||||
assert isinstance(result["content"], list)
|
||||
return string_value(object_value(result["content"][0])["text"])
|
||||
|
||||
|
||||
def _tool_call_authorizations(peer: McpPeer, seen: list[dict[str, object]]) -> tuple[object, ...]:
|
||||
seen.extend(peer.drain())
|
||||
return tuple(
|
||||
object_value(call["headers"]).get("authorization")
|
||||
for call in seen
|
||||
if isinstance(call["body"], dict) and call["body"].get("method") == "tools/call"
|
||||
)
|
||||
|
||||
|
||||
def _spend_rows(marker: Canary, since: datetime, call_types: frozenset[str]) -> Sequence[Mapping[str, object]]:
|
||||
"""Every spend row carrying ``marker``, once a row of each of ``call_types`` has been written."""
|
||||
return eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, call_type FROM "LiteLLM_SpendLogs" '
|
||||
'WHERE "startTime" >= %s AND proxy_server_request::text LIKE %s',
|
||||
(since.astimezone(UTC).replace(tzinfo=None) - SLACK, f"%{marker.core}%"),
|
||||
),
|
||||
lambda rows: call_types <= {row["call_type"] for row in rows},
|
||||
seconds=70,
|
||||
)
|
||||
|
||||
|
||||
def _drawer_hits(
|
||||
rig: Rig, request_ids: Sequence[str], canaries: Sequence[Canary], callers: Mapping[str, str]
|
||||
) -> tuple[Hit, ...]:
|
||||
"""S2 for the Logs drawer of every extra spend row the scenario wrote."""
|
||||
found: Final[list[Hit]] = [] # mutable-ok: accumulated across rows and callers
|
||||
for request_id in request_ids:
|
||||
for label, key in callers.items():
|
||||
response = rig.proxy.request("GET", f"/spend/logs/ui/{request_id}", key=key)
|
||||
where = f"GET /spend/logs/ui/{request_id} as {label} -> {response.status_code}"
|
||||
found.extend(
|
||||
Hit("S2", where, match.slot, match.encoding) for match in find_canary(response.content, canaries)
|
||||
)
|
||||
return tuple(found)
|
||||
|
||||
|
||||
@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~400 GET routes as two callers
|
||||
@pytest.mark.parametrize("outcome", ["success", "upstream_401"])
|
||||
@pytest.mark.parametrize("via", ["direct", "chat"])
|
||||
@pytest.mark.parametrize("slot", ["F1", "F2", "F2E", "F3"])
|
||||
def test_mcp_credential_reaches_only_its_peer(
|
||||
rig: Rig, slot: str, via: Via, outcome: Outcome, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
credential: Final = canary(slot)
|
||||
marker: Final = canary(MARKER)
|
||||
started: Final = datetime.now(UTC)
|
||||
peer_calls: Final[list[dict[str, object]]] = [] # mutable-ok: accumulates the peer's recorded requests
|
||||
with (
|
||||
scripted_peer(echo_tool("echo"), ScriptedTool("deny", _deny)) as peer,
|
||||
rig.proxy.scenario() as scenario,
|
||||
_wired(slot, rig, scenario, peer, credential) as wiring,
|
||||
):
|
||||
response: Final = _send(rig, wiring, via, TOOL[outcome], f"slot {slot} {marker.value}")
|
||||
answer: Final = _answer(response, via)
|
||||
assert (marker.value in answer) if outcome == "success" else ("401" in answer), answer
|
||||
assert _tool_call_authorizations(peer, peer_calls) == (f"Bearer {credential.value}",), (
|
||||
f"Positive control: the MCP peer never received the {slot} canary on tools/call"
|
||||
)
|
||||
rows: Final = _spend_rows(
|
||||
marker, started, frozenset({"call_mcp_tool", "acompletion"} if via == "chat" else {"call_mcp_tool"})
|
||||
)
|
||||
tool_row: Final = next(str(row["request_id"]) for row in rows if row["call_type"] == "call_mcp_tool")
|
||||
settle(rig, tool_row, marker)
|
||||
|
||||
canaries: Final = (marker, credential)
|
||||
report: Final = sweep_all(
|
||||
rig.proxy,
|
||||
canaries,
|
||||
responses=(response, *wiring.responses),
|
||||
sinks={name: sink.requests() for name, sink in rig.sinks.items()},
|
||||
ids={
|
||||
"request_id": tool_row,
|
||||
"server_id": wiring.server_id,
|
||||
# The OAuth discovery routes keyed by server name exist only for OAuth servers.
|
||||
**({"mcp_server_name": wiring.alias} if slot == "F2" else {}),
|
||||
"team_id": wiring.caller.team_id,
|
||||
"user_id": wiring.caller.user_id,
|
||||
"model": CONFIG_MODEL,
|
||||
"model_id": rig.model_id,
|
||||
},
|
||||
callers=wiring.caller.callers(rig),
|
||||
own_headers=rig.own_headers,
|
||||
since=started,
|
||||
)
|
||||
record_route_sweep(report.routes, request.node.nodeid)
|
||||
assert_marker_seen(report, {"S2": f"GET /spend/logs?request_id={tool_row} as admin -> 200"})
|
||||
assert_marker_seen(
|
||||
report,
|
||||
{
|
||||
"S1": "LiteLLM_SpendLogs.proxy_server_request",
|
||||
"S2": f"GET /spend/logs/ui/{tool_row} as admin -> 200",
|
||||
"S4": f"{GENERIC_SINK}[",
|
||||
**({"S3": "response[0] POST"} if outcome == "success" else {}),
|
||||
},
|
||||
)
|
||||
other_rows: Final = tuple(str(row["request_id"]) for row in rows if str(row["request_id"]) != tool_row)
|
||||
assert_no_hits(
|
||||
(*report.credential_hits(), *_drawer_hits(rig, other_rows, (credential,), wiring.caller.callers(rig))),
|
||||
f"slot {slot}, {via}, {outcome}",
|
||||
)
|
||||
210
tests/integration/security/test_passthrough_slots.py
Normal file
210
tests/integration/security/test_passthrough_slots.py
Normal file
|
|
@ -0,0 +1,210 @@
|
|||
"""Slots H1 and H2: pass-through, vector store and search tool credentials reach only their upstream.
|
||||
|
||||
Each test boots an owned proxy whose config declares all three credentials against one
|
||||
recording upstream double:
|
||||
|
||||
- H1: a pass-through endpoint whose ``Authorization`` header is ``Bearer os.environ/<name>``,
|
||||
with the canary in that environment variable;
|
||||
- H2: an OpenAI vector store in ``vector_store_registry`` with the canary as ``api_key``;
|
||||
- H2S: a Perplexity search tool in ``search_tools`` with the canary as ``api_key``.
|
||||
|
||||
The test sends one request through the slot's route, and the upstream answers 200 or, when the
|
||||
request carries ``UPSTREAM_REJECT``, 401. Positive control: the upstream must receive
|
||||
``Authorization: Bearer <canary>`` on the request carrying the marker, or the test fails before
|
||||
sweeping. Sensitivity control: the marker must be reported where stored prompts belong. Then no
|
||||
sweep may find any of the three canaries.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Scenario, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from integration.security._canary import MARKER, Canary, canary
|
||||
from integration.security._sinks import CONFIG_MODEL, GENERIC_SINK, Caller, Recorder, Rig, canary_rig, settle
|
||||
from integration.security._sweeps import assert_marker_seen, assert_no_hits, record_route_sweep, sweep_all
|
||||
|
||||
Outcome = Literal["success", "upstream_401"]
|
||||
PASS_THROUGH_ROUTE: Final = "/canary-pass-through"
|
||||
PASS_THROUGH_ENV: Final = "CANARY_PASS_THROUGH_KEY"
|
||||
VECTOR_STORE_ID: Final = "canary-vector-store"
|
||||
SEARCH_TOOL: Final = "canary-search-tool"
|
||||
UPSTREAM_REJECT: Final = "canary-upstream-reject"
|
||||
SLOTS: Final = ("H1", "H2", "H2S")
|
||||
SLACK: Final = timedelta(seconds=5)
|
||||
|
||||
|
||||
def _upstream(request: Request) -> Reply:
|
||||
"""Pass-through, OpenAI vector store search and Perplexity search double."""
|
||||
if UPSTREAM_REJECT.encode() in request.body:
|
||||
return Reply(status=401, body=b'{"error":"invalid credentials"}')
|
||||
body: Final = json.loads(request.body or b"{}")
|
||||
query: Final = str(body.get("query", ""))
|
||||
if request.target.startswith("/v1/vector_stores/"):
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"object": "vector_store.search_results.page",
|
||||
"search_query": [query],
|
||||
"data": [
|
||||
{
|
||||
"file_id": "file-canary",
|
||||
"filename": "canary.txt",
|
||||
"score": 0.9,
|
||||
"attributes": {},
|
||||
"content": [{"type": "text", "text": query}],
|
||||
}
|
||||
],
|
||||
"has_more": False,
|
||||
"next_page": None,
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
if request.target == "/search":
|
||||
return Reply(
|
||||
body=json.dumps({"results": [{"title": "canary", "url": "https://example.com", "snippet": query}]}).encode()
|
||||
)
|
||||
return Reply(body=json.dumps({"received": body}).encode())
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Upstreamed:
|
||||
rig: Rig
|
||||
upstream: Recorder
|
||||
canaries: Mapping[str, Canary]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rigged(tmp_path: Path) -> Iterator[Upstreamed]:
|
||||
"""One owned proxy per test: the H credentials live in its config and environment."""
|
||||
canaries: Final = {slot: canary(slot) for slot in SLOTS}
|
||||
with wire_server(_upstream) as wire:
|
||||
|
||||
def configure(config: dict[str, object], provider_url: str) -> None:
|
||||
general: Final = config["general_settings"]
|
||||
assert isinstance(general, dict)
|
||||
general["pass_through_endpoints"] = [
|
||||
{
|
||||
"path": PASS_THROUGH_ROUTE,
|
||||
"target": wire.url + "/pass-through",
|
||||
"headers": {"Authorization": f"Bearer os.environ/{PASS_THROUGH_ENV}"},
|
||||
"auth": True,
|
||||
}
|
||||
]
|
||||
config["vector_store_registry"] = [
|
||||
{
|
||||
"vector_store_name": VECTOR_STORE_ID,
|
||||
"litellm_params": {
|
||||
"vector_store_id": VECTOR_STORE_ID,
|
||||
"custom_llm_provider": "openai",
|
||||
"api_key": canaries["H2"].value,
|
||||
"api_base": wire.url + "/v1",
|
||||
},
|
||||
}
|
||||
]
|
||||
config["search_tools"] = [
|
||||
{
|
||||
"search_tool_name": SEARCH_TOOL,
|
||||
"litellm_params": {
|
||||
"search_provider": "perplexity",
|
||||
"api_key": canaries["H2S"].value,
|
||||
"api_base": wire.url,
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
with canary_rig(tmp_path, configure=configure, environment={PASS_THROUGH_ENV: canaries["H1"].value}) as rig:
|
||||
yield Upstreamed(rig, Recorder(wire), canaries)
|
||||
|
||||
|
||||
def _caller(scenario: Scenario) -> Caller:
|
||||
team: Final = scenario.team(metadata={"allowed_passthrough_routes": [PASS_THROUGH_ROUTE]})
|
||||
user: Final = scenario.user(user_role="internal_user")
|
||||
scenario.gateway.post("/team/member_add", {"team_id": team, "member": {"user_id": user, "role": "user"}})
|
||||
key: Final = scenario.key(team_id=team, user_id=user, models=[CONFIG_MODEL])
|
||||
return Caller(team, user, key)
|
||||
|
||||
|
||||
def _send(rig: Rig, slot: str, key: str, text: str) -> httpx.Response:
|
||||
if slot == "H1":
|
||||
return rig.proxy.request("POST", PASS_THROUGH_ROUTE, {"text": text}, key=key)
|
||||
if slot == "H2":
|
||||
return rig.proxy.request("POST", f"/v1/vector_stores/{VECTOR_STORE_ID}/search", {"query": text}, key=key)
|
||||
return rig.proxy.request("POST", f"/v1/search/{SEARCH_TOOL}", {"query": text}, key=key)
|
||||
|
||||
|
||||
def _spend_row(marker: Canary, since: datetime) -> str:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE "startTime" >= %s AND proxy_server_request::text LIKE %s',
|
||||
(since.astimezone(UTC).replace(tzinfo=None) - SLACK, f"%{marker.core}%"),
|
||||
),
|
||||
lambda found: len(found) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return str(rows[0]["request_id"])
|
||||
|
||||
|
||||
@pytest.mark.timeout(240) # full S1/S2 walk: every table and ~400 GET routes as two callers
|
||||
@pytest.mark.parametrize("outcome", ["success", "upstream_401"])
|
||||
@pytest.mark.parametrize("slot", SLOTS)
|
||||
def test_upstream_credential_reaches_only_its_upstream(
|
||||
rigged: Upstreamed, slot: str, outcome: Outcome, request: pytest.FixtureRequest
|
||||
) -> None:
|
||||
rig: Final = rigged.rig
|
||||
credential: Final = rigged.canaries[slot]
|
||||
marker: Final = canary(MARKER)
|
||||
text: Final = f"slot {slot} {marker.value}" + (f" {UPSTREAM_REJECT}" if outcome == "upstream_401" else "")
|
||||
started: Final = datetime.now(UTC)
|
||||
with rig.proxy.scenario() as scenario:
|
||||
caller: Final = _caller(scenario)
|
||||
response: Final = _send(rig, slot, caller.key, text)
|
||||
assert response.status_code == (200 if outcome == "success" else 401), response.text
|
||||
delivered: Final = rigged.upstream.carrying(marker.core)
|
||||
assert [received.headers.get("authorization") for received in delivered] == [f"Bearer {credential.value}"], (
|
||||
f"Positive control: the upstream never received the {slot} canary"
|
||||
)
|
||||
request_id: Final = _spend_row(marker, started)
|
||||
delivers_to_sink: Final = not (slot == "H1" and outcome == "upstream_401")
|
||||
if delivers_to_sink:
|
||||
settle(rig, request_id, marker)
|
||||
|
||||
report: Final = sweep_all(
|
||||
rig.proxy,
|
||||
(marker, *rigged.canaries.values()),
|
||||
responses=(response,),
|
||||
sinks={name: sink.requests() for name, sink in rig.sinks.items()},
|
||||
ids={
|
||||
"request_id": request_id,
|
||||
"team_id": caller.team_id,
|
||||
"user_id": caller.user_id,
|
||||
"vector_store_id": VECTOR_STORE_ID,
|
||||
"search_tool_name": SEARCH_TOOL,
|
||||
"model": CONFIG_MODEL,
|
||||
"model_id": rig.model_id,
|
||||
},
|
||||
callers=caller.callers(rig),
|
||||
own_headers=rig.own_headers,
|
||||
since=started,
|
||||
)
|
||||
record_route_sweep(report.routes, request.node.nodeid)
|
||||
assert_marker_seen(report, {"S2": f"GET /spend/logs?request_id={request_id} as admin -> 200"})
|
||||
assert_marker_seen(
|
||||
report,
|
||||
{
|
||||
"S1": "LiteLLM_SpendLogs.proxy_server_request",
|
||||
"S2": f"GET /spend/logs/ui/{request_id} as admin -> 200",
|
||||
**({"S3": "response[0] POST"} if outcome == "success" else {}),
|
||||
**({"S4": f"{GENERIC_SINK}["} if delivers_to_sink else {}),
|
||||
},
|
||||
)
|
||||
assert_no_hits(report.credential_hits(), f"slot {slot}, {outcome}")
|
||||
Loading…
Add table
Reference in a new issue