mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(chatgpt,github_copilot): refuse device-code login inside an event loop or worker thread (#39585)
* fix(chatgpt,github_copilot): refuse device-code login when an event loop is running Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(github_copilot): drop stray whitespace change Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(chatgpt): bound token refresh timeout and drop placeholder assignment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(chatgpt,github_copilot): keep the token file path out of the event-loop 401 message * fix(auth): refuse device-code login from worker threads too /v1/messages runs its handler in an executor thread, where the running-loop check never fires, so a chatgpt or github_copilot model still started the interactive device-code login there and the request hung for up to 15 minutes. The guard now also requires the main thread, so the login only runs where a human can actually answer it. * test(chatgpt): keep authenticator tests out of the real token directory * fix(chatgpt): keep the 5 second connect timeout and the operator's request_timeout on the token refresh call * test(integration): cover the device-code login guard on the proxy and the SDK --------- Co-authored-by: mateo <mateo@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
64cddd6e13
commit
ab3a59fe81
9 changed files with 2128 additions and 8 deletions
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import functools
|
||||
import threading
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -68,6 +69,18 @@ def asyncify(
|
|||
return wrapper
|
||||
|
||||
|
||||
def is_event_loop_running() -> bool:
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def can_block_current_thread() -> bool:
|
||||
return threading.current_thread() is threading.main_thread() and not is_event_loop_running()
|
||||
|
||||
|
||||
def run_async_function(async_function, *args, **kwargs):
|
||||
"""
|
||||
Helper utility to run an async function in a sync context.
|
||||
|
|
|
|||
|
|
@ -9,6 +9,9 @@ import httpx
|
|||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS
|
||||
from litellm.litellm_core_utils.asyncify import can_block_current_thread
|
||||
from litellm.litellm_core_utils.request_timeout_resolver import get_configured_request_timeout
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
|
||||
from .common_utils import (
|
||||
|
|
@ -25,6 +28,7 @@ from .common_utils import (
|
|||
)
|
||||
|
||||
TOKEN_EXPIRY_SKEW_SECONDS: Final = 60
|
||||
TOKEN_REFRESH_TIMEOUT_SECONDS: Final = 30
|
||||
DEVICE_CODE_TIMEOUT_SECONDS: Final = 15 * 60
|
||||
DEVICE_CODE_COOLDOWN_SECONDS: Final = 5 * 60
|
||||
DEVICE_CODE_POLL_SLEEP_SECONDS: Final = 5
|
||||
|
|
@ -40,6 +44,14 @@ def _optional_str(value: JsonValue | None) -> str | None:
|
|||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
def _token_refresh_timeout() -> httpx.Timeout:
|
||||
configured: Final = get_configured_request_timeout()
|
||||
seconds: Final = (
|
||||
TOKEN_REFRESH_TIMEOUT_SECONDS if configured is None else min(configured, TOKEN_REFRESH_TIMEOUT_SECONDS)
|
||||
)
|
||||
return httpx.Timeout(seconds, connect=HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS)
|
||||
|
||||
|
||||
class Authenticator:
|
||||
def __init__(self) -> None:
|
||||
self.token_dir = os.getenv(
|
||||
|
|
@ -66,6 +78,18 @@ class Authenticator:
|
|||
except RefreshAccessTokenError as exc:
|
||||
verbose_logger.warning("ChatGPT refresh token failed, re-login required: %s", exc)
|
||||
|
||||
if not can_block_current_thread():
|
||||
raise GetAccessTokenError(
|
||||
message=(
|
||||
"ChatGPT device-code login needs a human and cannot run inside a running event loop "
|
||||
"or a worker thread (for example the LiteLLM proxy). Log in once outside the proxy with "
|
||||
'`python -c "from litellm.llms.chatgpt.authenticator import Authenticator; '
|
||||
'Authenticator().get_access_token()"` and mount the resulting auth.json into the proxy, '
|
||||
"or set CHATGPT_TOKEN_DIR to a directory that already holds it."
|
||||
),
|
||||
status_code=401,
|
||||
)
|
||||
|
||||
cooldown_remaining: Final = self._get_device_code_cooldown_remaining(auth_data)
|
||||
if cooldown_remaining > 0:
|
||||
token: Final = self._wait_for_access_token(cooldown_remaining)
|
||||
|
|
@ -309,6 +333,7 @@ class Authenticator:
|
|||
"refresh_token": refresh_token,
|
||||
"scope": "openid profile email",
|
||||
},
|
||||
timeout=_token_refresh_timeout(),
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data: Final = _JSON_OBJECT_ADAPTER.validate_python(resp.json())
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from typing import Any, Final
|
|||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.asyncify import can_block_current_thread
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
|
||||
from .common_utils import (
|
||||
|
|
@ -57,6 +58,18 @@ class Authenticator:
|
|||
except OSError:
|
||||
verbose_logger.warning("No existing access token found or error reading file")
|
||||
|
||||
if not can_block_current_thread():
|
||||
raise GetAccessTokenError(
|
||||
message=(
|
||||
"GitHub Copilot device-code login needs a human and cannot run inside a running event loop "
|
||||
"or a worker thread (for example the LiteLLM proxy). Log in once outside the proxy with "
|
||||
'`python -c "from litellm.llms.github_copilot.authenticator import Authenticator; '
|
||||
'Authenticator().get_access_token()"` and mount the resulting access-token file into '
|
||||
"the proxy, or set GITHUB_COPILOT_TOKEN_DIR to a directory that already holds it."
|
||||
),
|
||||
status_code=401,
|
||||
)
|
||||
|
||||
for attempt in range(3):
|
||||
verbose_logger.debug("Access token acquisition attempt %s/3", attempt + 1)
|
||||
try:
|
||||
|
|
|
|||
383
tests/integration/_support/device_login.py
Normal file
383
tests/integration/_support/device_login.py
Normal file
|
|
@ -0,0 +1,383 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import contextlib
|
||||
import io
|
||||
import json
|
||||
import socket
|
||||
import socketserver
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Generator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
from integration._support import responses_vendor as rv
|
||||
from integration._support.tls import server_context, write_self_signed_cert
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
CHATGPT_AUTH_HOST: Final = "auth.openai.com"
|
||||
GITHUB_HOST: Final = "github.com"
|
||||
GITHUB_API_HOST: Final = "api.github.com"
|
||||
AUTH_HOSTS: Final = (CHATGPT_AUTH_HOST, GITHUB_HOST, GITHUB_API_HOST)
|
||||
CHATGPT_ACCOUNT: Final = "acct-device-login"
|
||||
_FAR_FUTURE: Final = 4102444800
|
||||
_LONG_AGO: Final = 946684800
|
||||
_OPENAI_AUTH_CLAIM: Final = "https://api.openai.com/auth"
|
||||
|
||||
CHATGPT_REFUSAL: Final = (
|
||||
"ChatGPT device-code login needs a human and cannot run inside a running event loop "
|
||||
"or a worker thread (for example the LiteLLM proxy). Log in once outside the proxy with "
|
||||
'`python -c "from litellm.llms.chatgpt.authenticator import Authenticator; '
|
||||
'Authenticator().get_access_token()"` and mount the resulting auth.json into the proxy, '
|
||||
"or set CHATGPT_TOKEN_DIR to a directory that already holds it."
|
||||
)
|
||||
COPILOT_REFUSAL: Final = (
|
||||
"GitHub Copilot device-code login needs a human and cannot run inside a running event loop "
|
||||
"or a worker thread (for example the LiteLLM proxy). Log in once outside the proxy with "
|
||||
'`python -c "from litellm.llms.github_copilot.authenticator import Authenticator; '
|
||||
'Authenticator().get_access_token()"` and mount the resulting access-token file into '
|
||||
"the proxy, or set GITHUB_COPILOT_TOKEN_DIR to a directory that already holds it."
|
||||
)
|
||||
|
||||
|
||||
def _segment(value: Mapping[str, JsonValue]) -> str:
|
||||
return base64.urlsafe_b64encode(json.dumps(value).encode()).rstrip(b"=").decode()
|
||||
|
||||
|
||||
def chatgpt_jwt(subject: str, expires_at: int = _FAR_FUTURE) -> str:
|
||||
claims: Final[Mapping[str, JsonValue]] = {
|
||||
"sub": subject,
|
||||
"exp": expires_at,
|
||||
_OPENAI_AUTH_CLAIM: {"chatgpt_account_id": CHATGPT_ACCOUNT},
|
||||
}
|
||||
return f"{_segment({'alg': 'none', 'typ': 'JWT'})}.{_segment(claims)}.synthetic"
|
||||
|
||||
|
||||
CHATGPT_STORED: Final = chatgpt_jwt("stored")
|
||||
CHATGPT_REFRESHED: Final = chatgpt_jwt("refreshed")
|
||||
CHATGPT_FIRST_LOGIN: Final = chatgpt_jwt("first-login")
|
||||
CHATGPT_REJECTED: Final = chatgpt_jwt("rejected")
|
||||
CHATGPT_EXPIRED: Final = chatgpt_jwt("expired", _LONG_AGO)
|
||||
CHATGPT_ID_TOKEN: Final = chatgpt_jwt("identity")
|
||||
GOOD_REFRESH: Final = "device-login-refresh-good"
|
||||
REVOKED_REFRESH: Final = "device-login-refresh-revoked"
|
||||
CHATGPT_USER_CODE: Final = "CGPT-LEAK"
|
||||
COPILOT_USER_CODE: Final = "COPI-LEAK"
|
||||
COPILOT_ACCESS: Final = "device-login-github-access"
|
||||
COPILOT_REJECTED_ACCESS: Final = "device-login-github-access-rejected"
|
||||
COPILOT_FIRST_LOGIN_ACCESS: Final = "device-login-github-access-first-login"
|
||||
COPILOT_STORED_KEY: Final = "tid=device-login;stored"
|
||||
COPILOT_MINTED_KEY: Final = "tid=device-login;minted"
|
||||
COPILOT_REJECTED_KEY: Final = "tid=device-login;rejected"
|
||||
CONTROL_KEY: Final = "device-login-control-key"
|
||||
ACCEPTED_BEARERS: Final = frozenset(
|
||||
{CHATGPT_STORED, CHATGPT_REFRESHED, CHATGPT_FIRST_LOGIN, COPILOT_STORED_KEY, COPILOT_MINTED_KEY, CONTROL_KEY}
|
||||
)
|
||||
REJECTED_DETAIL: Final = "the scripted provider rejected this bearer token"
|
||||
_DEVICE_AUTH_ID: Final = "device-auth-synthetic"
|
||||
_AUTHORIZATION_CODE: Final = "authorization-code-synthetic"
|
||||
_CODE_VERIFIER: Final = "code-verifier-synthetic"
|
||||
_GITHUB_DEVICE_CODE: Final = "github-device-code-synthetic"
|
||||
_EMBEDDING: Final[Mapping[str, JsonValue]] = MappingProxyType(
|
||||
{
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.25, 0.5, 0.75]}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 3, "total_tokens": 3},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _json(status: int, body: Mapping[str, JsonValue]) -> Reply:
|
||||
return Reply(status=status, body=json.dumps(dict(body)).encode())
|
||||
|
||||
|
||||
def bearer(request: Request) -> str:
|
||||
return request.headers.get("authorization", "").removeprefix("Bearer ")
|
||||
|
||||
|
||||
def _provider_api(vendor: rv.ResponsesVendor) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
if bearer(request) not in ACCEPTED_BEARERS:
|
||||
return _json(401, {"detail": REJECTED_DETAIL})
|
||||
if urlsplit(request.target).path.endswith("/embeddings"):
|
||||
return _json(200, _EMBEDDING)
|
||||
return vendor.respond(request)
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _chatgpt_tokens(access_token: str) -> Mapping[str, JsonValue]:
|
||||
return {"access_token": access_token, "id_token": CHATGPT_ID_TOKEN, "refresh_token": GOOD_REFRESH}
|
||||
|
||||
|
||||
def _oauth_token(request: Request, grant: threading.Event) -> Reply:
|
||||
if request.headers.get("content-type", "").startswith("application/x-www-form-urlencoded"):
|
||||
form: Final = parse_qs(request.body.decode())
|
||||
exchanged: Final = form.get("code") == [_AUTHORIZATION_CODE] and form.get("code_verifier") == [_CODE_VERIFIER]
|
||||
if grant.is_set() and exchanged:
|
||||
return _json(200, _chatgpt_tokens(CHATGPT_FIRST_LOGIN))
|
||||
return _json(400, {"error": "invalid_grant"})
|
||||
body: Final = rv.JSON_OBJECT.validate_json(request.body)
|
||||
if body.get("grant_type") == "refresh_token" and body.get("refresh_token") == GOOD_REFRESH:
|
||||
return _json(200, _chatgpt_tokens(CHATGPT_REFRESHED))
|
||||
return _json(400, {"error": "invalid_grant", "error_description": "refresh token was revoked"})
|
||||
|
||||
|
||||
def _copilot_api_key(request: Request, api_url: str) -> Reply:
|
||||
if request.headers.get("authorization") not in (f"token {COPILOT_ACCESS}", f"token {COPILOT_FIRST_LOGIN_ACCESS}"):
|
||||
return _json(401, {"message": "Bad credentials"})
|
||||
return _json(
|
||||
200,
|
||||
{"token": COPILOT_MINTED_KEY, "expires_at": int(time.time()) + 1800, "endpoints": {"api": api_url}},
|
||||
)
|
||||
|
||||
|
||||
def _switched_off() -> Reply:
|
||||
return _json(503, {"error": "device login is switched off in this scenario"})
|
||||
|
||||
|
||||
def _auth_hosts(api_url: str, grant: threading.Event) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
match (request.headers.get("host", ""), request.method, urlsplit(request.target).path):
|
||||
case ("auth.openai.com", "POST", "/oauth/token"):
|
||||
return _oauth_token(request, grant)
|
||||
case ("auth.openai.com", "POST", "/api/accounts/deviceauth/usercode") if grant.is_set():
|
||||
return _json(200, {"device_auth_id": _DEVICE_AUTH_ID, "user_code": CHATGPT_USER_CODE, "interval": "5"})
|
||||
case ("auth.openai.com", "POST", "/api/accounts/deviceauth/token") if grant.is_set():
|
||||
return _json(
|
||||
200,
|
||||
{
|
||||
"authorization_code": _AUTHORIZATION_CODE,
|
||||
"code_challenge": "code-challenge-synthetic",
|
||||
"code_verifier": _CODE_VERIFIER,
|
||||
},
|
||||
)
|
||||
case ("github.com", "POST", "/login/device/code") if grant.is_set():
|
||||
return _json(
|
||||
200,
|
||||
{
|
||||
"device_code": _GITHUB_DEVICE_CODE,
|
||||
"user_code": COPILOT_USER_CODE,
|
||||
"verification_uri": "https://github.com/login/device",
|
||||
"expires_in": 900,
|
||||
"interval": 5,
|
||||
},
|
||||
)
|
||||
case ("github.com", "POST", "/login/oauth/access_token") if grant.is_set():
|
||||
return _json(200, {"access_token": COPILOT_FIRST_LOGIN_ACCESS, "token_type": "bearer"})
|
||||
case ("api.github.com", "GET", "/copilot_internal/v2/token"):
|
||||
return _copilot_api_key(request, api_url)
|
||||
case _:
|
||||
return _switched_off()
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Hangup:
|
||||
authority: str
|
||||
seconds: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Switches:
|
||||
grant: threading.Event
|
||||
hold_connect: threading.Event
|
||||
stall_tls: threading.Event
|
||||
refuse_connect: threading.Event
|
||||
|
||||
def clear(self) -> None:
|
||||
for switch in (self.grant, self.hold_connect, self.stall_tls, self.refuse_connect):
|
||||
switch.clear()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Peers:
|
||||
api: Wire
|
||||
auth: Wire
|
||||
tunnel: str
|
||||
cert: Path
|
||||
authorities: SimpleQueue[str]
|
||||
hangups: SimpleQueue[Hangup]
|
||||
switches: Switches
|
||||
|
||||
def environment(self) -> Mapping[str, str]:
|
||||
return MappingProxyType(
|
||||
{
|
||||
"HTTPS_PROXY": self.tunnel,
|
||||
"NO_PROXY": "127.0.0.1,localhost",
|
||||
"SSL_CERT_FILE": str(self.cert),
|
||||
"CHATGPT_API_BASE": self.api.url,
|
||||
"GITHUB_COPILOT_API_BASE": self.api.url,
|
||||
}
|
||||
)
|
||||
|
||||
def auth_connections(self) -> tuple[str, ...]:
|
||||
return tuple(self.authorities.get_nowait() for _ in range(self.authorities.qsize()))
|
||||
|
||||
def dropped(self) -> tuple[Hangup, ...]:
|
||||
return tuple(self.hangups.get_nowait() for _ in range(self.hangups.qsize()))
|
||||
|
||||
def reset(self) -> None:
|
||||
self.switches.clear()
|
||||
self.api.drain()
|
||||
self.auth.drain()
|
||||
self.auth_connections()
|
||||
self.dropped()
|
||||
|
||||
|
||||
def _pipe(source: socket.socket, sink: socket.socket) -> None:
|
||||
with contextlib.suppress(OSError):
|
||||
for chunk in iter(lambda: source.recv(65536), b""):
|
||||
sink.sendall(chunk)
|
||||
with contextlib.suppress(OSError):
|
||||
sink.shutdown(socket.SHUT_WR)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Held:
|
||||
outcome: Literal["released", "hung-up", "cut"]
|
||||
early_bytes: bytes
|
||||
|
||||
|
||||
def _hold(connection: socket.socket, held: threading.Event, cut: threading.Event) -> _Held:
|
||||
received: Final = io.BytesIO()
|
||||
connection.settimeout(0.1)
|
||||
while held.is_set() and not cut.is_set():
|
||||
try:
|
||||
chunk: Final = connection.recv(65536)
|
||||
except TimeoutError:
|
||||
continue
|
||||
except OSError:
|
||||
return _Held("hung-up", b"")
|
||||
if chunk == b"":
|
||||
return _Held("hung-up", b"")
|
||||
received.write(chunk)
|
||||
if cut.is_set():
|
||||
return _Held("cut", b"")
|
||||
return _Held("released", received.getvalue())
|
||||
|
||||
|
||||
class _TunnelServer(socketserver.ThreadingTCPServer):
|
||||
daemon_threads = True
|
||||
allow_reuse_address = True
|
||||
request_queue_size = 128
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _tunnel(
|
||||
destination: Wire, switches: Switches, authorities: SimpleQueue[str], hangups: SimpleQueue[Hangup]
|
||||
) -> Generator[str]:
|
||||
destination_port: Final = int(destination.url.rsplit(":", 1)[1])
|
||||
|
||||
class Tunnel(socketserver.StreamRequestHandler):
|
||||
rbufsize = 0
|
||||
request: socket.socket
|
||||
|
||||
def hold(self, authority: str, opened: float, switch: threading.Event) -> bytes | None:
|
||||
held: Final = _hold(self.request, switch, switches.refuse_connect)
|
||||
match held.outcome:
|
||||
case "released":
|
||||
return held.early_bytes
|
||||
case "hung-up":
|
||||
hangups.put(Hangup(authority, time.monotonic() - opened))
|
||||
return None
|
||||
case "cut":
|
||||
return None
|
||||
|
||||
def handle(self) -> None:
|
||||
opened: Final = time.monotonic()
|
||||
request_line: Final = self.rfile.readline().decode().split()
|
||||
while self.rfile.readline() not in (b"\r\n", b""):
|
||||
pass
|
||||
if len(request_line) < 2:
|
||||
return
|
||||
authority: Final = request_line[1]
|
||||
if authority.rsplit(":", 1)[0] not in AUTH_HOSTS:
|
||||
self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n")
|
||||
return
|
||||
authorities.put(authority)
|
||||
if switches.refuse_connect.is_set():
|
||||
self.wfile.write(b"HTTP/1.1 502 Bad Gateway\r\ncontent-length: 0\r\n\r\n")
|
||||
return
|
||||
if self.hold(authority, opened, switches.hold_connect) is None:
|
||||
return
|
||||
self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n")
|
||||
client_hello: Final = self.hold(authority, opened, switches.stall_tls)
|
||||
if client_hello is None:
|
||||
return
|
||||
self.request.settimeout(10)
|
||||
with socket.create_connection(("127.0.0.1", destination_port), timeout=10) as upstream:
|
||||
upstream.sendall(client_hello)
|
||||
outbound: Final = threading.Thread(target=_pipe, args=(self.request, upstream))
|
||||
outbound.start()
|
||||
_pipe(upstream, self.request)
|
||||
outbound.join(timeout=12)
|
||||
|
||||
with _TunnelServer(("127.0.0.1", 0), Tunnel) as server:
|
||||
thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05})
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_address[1]}"
|
||||
finally:
|
||||
switches.clear()
|
||||
server.shutdown()
|
||||
thread.join(timeout=6)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def device_login_peers(directory: Path) -> Generator[Peers]:
|
||||
cert, key = write_self_signed_cert(directory, AUTH_HOSTS)
|
||||
switches: Final = Switches(threading.Event(), threading.Event(), threading.Event(), threading.Event())
|
||||
authorities: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
hangups: Final[SimpleQueue[Hangup]] = SimpleQueue()
|
||||
with (
|
||||
wire_server(_provider_api(rv.ResponsesVendor())) as api,
|
||||
wire_server(_auth_hosts(api.url, switches.grant), tls=server_context(cert, key)) as auth,
|
||||
_tunnel(auth, switches, authorities, hangups) as tunnel,
|
||||
):
|
||||
yield Peers(api, auth, tunnel, cert, authorities, hangups, switches)
|
||||
|
||||
|
||||
def chatgpt_record(
|
||||
access_token: str, *, refresh_token: str | None = None, expires_at: JsonValue = None
|
||||
) -> Mapping[str, JsonValue]:
|
||||
optional: Final[Mapping[str, JsonValue]] = MappingProxyType(
|
||||
{"refresh_token": refresh_token, "expires_at": expires_at}
|
||||
)
|
||||
return MappingProxyType(
|
||||
{
|
||||
"access_token": access_token,
|
||||
"id_token": CHATGPT_ID_TOKEN,
|
||||
"account_id": CHATGPT_ACCOUNT,
|
||||
**{key: value for key, value in optional.items() if value is not None},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def write_chatgpt(directory: Path, record: Mapping[str, JsonValue]) -> None:
|
||||
(directory / "auth.json").write_text(json.dumps(dict(record)))
|
||||
|
||||
|
||||
def write_copilot_key(directory: Path, token: str, api_url: str, expires_in: float = 3600) -> None:
|
||||
(directory / "api-key.json").write_text(
|
||||
json.dumps({"token": token, "expires_at": time.time() + expires_in, "endpoints": {"api": api_url}})
|
||||
)
|
||||
|
||||
|
||||
def write_copilot_access(directory: Path, access_token: str) -> None:
|
||||
(directory / "access-token").write_text(access_token)
|
||||
|
||||
|
||||
def clear_tokens(*directories: Path) -> None:
|
||||
for directory in directories:
|
||||
for path in directory.iterdir():
|
||||
path.unlink()
|
||||
435
tests/integration/providers/test_device_code_login_guard_boot.py
Normal file
435
tests/integration/providers/test_device_code_login_guard_boot.py
Normal file
|
|
@ -0,0 +1,435 @@
|
|||
import asyncio
|
||||
import hashlib
|
||||
import itertools
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
import httpcore
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support import device_login as dl
|
||||
from integration._support import responses_vendor as rv
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(900)
|
||||
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_DROPPED: Final = re.compile(r"original model: (\S+), ignoring and continuing")
|
||||
_FIXED: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{"chatgpt-fixed": "chatgpt/gpt-5.6-terra", "copilot-fixed": "github_copilot/gpt-5.2"}
|
||||
)
|
||||
_WILDCARDS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{"chatgpt/*": "chatgpt/gpt-5.6-terra", "github_copilot/*": "github_copilot/gpt-5.2"}
|
||||
)
|
||||
_LOGIN_PREFIXES: Final = ("chatgpt/", "github_copilot/")
|
||||
_LOGIN_NAMES: Final = (*_FIXED, *_FIXED.values())
|
||||
_CONTROL_CALLS_BEFORE_THE_RELOAD: Final = 1
|
||||
_SHAPES: Final[Mapping[str, bool]] = MappingProxyType(
|
||||
{"docs-shapes": False, "docs-shapes-beside-a-universal-wildcard": True}
|
||||
)
|
||||
_NOT_FOUND: Final = "Invalid model name passed in model="
|
||||
_NO_HEALTHY: Final = "There are no healthy deployments for this model"
|
||||
_CONTROL: Final = "device-login-boot-control"
|
||||
_EXTRA: Final[Mapping[str, JsonValue]] = MappingProxyType({"num_retries": 0, "cache": {"no-cache": True}})
|
||||
_CALL_ID: Final = "x-litellm-call-id"
|
||||
_REQUEST_TIMEOUT_SECONDS: Final = 6
|
||||
_POLL_SECONDS: Final = 3
|
||||
_ADDRESS: Final = TypeAdapter(tuple[str, int])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
status: int
|
||||
text: str
|
||||
call_id: str
|
||||
|
||||
def message(self) -> str:
|
||||
error: Final = rv.JSON_OBJECT.validate_json(self.text)["error"]
|
||||
return str(rv.JSON_OBJECT.validate_python(error)["message"]) if isinstance(error, dict) else str(error)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _TokenDirs:
|
||||
chatgpt: Path
|
||||
copilot: Path
|
||||
|
||||
def environment(self) -> Mapping[str, str]:
|
||||
return MappingProxyType({"CHATGPT_TOKEN_DIR": str(self.chatgpt), "GITHUB_COPILOT_TOKEN_DIR": str(self.copilot)})
|
||||
|
||||
def store_valid_tokens(self, api_url: str) -> None:
|
||||
dl.write_chatgpt(self.chatgpt, dl.chatgpt_record(dl.CHATGPT_STORED))
|
||||
dl.write_copilot_key(self.copilot, dl.COPILOT_STORED_KEY, api_url)
|
||||
|
||||
|
||||
def _token_dirs(directory: Path) -> _TokenDirs:
|
||||
dirs: Final = _TokenDirs(directory / "chatgpt", directory / "copilot")
|
||||
dirs.chatgpt.mkdir()
|
||||
dirs.copilot.mkdir()
|
||||
return dirs
|
||||
|
||||
|
||||
def _config(control_api_url: str, directory: Path, *, fixed: bool, wildcards: bool, universal: bool = False) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
fixed_entries: Final = [{"model_name": name, "litellm_params": {"model": model}} for name, model in _FIXED.items()]
|
||||
wildcard_entries: Final = [{"model_name": pattern, "litellm_params": {"model": pattern}} for pattern in _WILDCARDS]
|
||||
config["model_list"] = [
|
||||
*(fixed_entries if fixed else []),
|
||||
*(wildcard_entries if wildcards else []),
|
||||
*([{"model_name": "*", "litellm_params": {"model": "*"}}] if universal else []),
|
||||
{
|
||||
"model_name": _CONTROL,
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": f"{control_api_url}/v1",
|
||||
"api_key": dl.CONTROL_KEY,
|
||||
},
|
||||
},
|
||||
]
|
||||
path: Final = directory / "device-login-boot.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _chat(base_url: str, key: str, model: str, marker: str, seconds: float = 60) -> _Served:
|
||||
with httpx.Client(base_url=base_url, timeout=seconds, trust_env=False) as client:
|
||||
response: Final = client.post(
|
||||
"/v1/chat/completions",
|
||||
json={"model": model, "messages": [{"role": "user", "content": f"Reply to marker-{marker}"}], **_EXTRA},
|
||||
headers={"Authorization": f"Bearer {key}"},
|
||||
)
|
||||
return _Served(response.status_code, response.text, response.headers.get(_CALL_ID, ""))
|
||||
|
||||
|
||||
def _model_ids(text: str) -> tuple[str, ...]:
|
||||
return tuple(
|
||||
sorted(str(item["id"]) for item in rv.ITEMS.validate_python(rv.JSON_OBJECT.validate_json(text)["data"]))
|
||||
)
|
||||
|
||||
|
||||
def _workers(owned: OwnedProxy) -> tuple[int, ...]:
|
||||
return eventually(
|
||||
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
|
||||
lambda pids: len(pids) == 2,
|
||||
seconds=30,
|
||||
)
|
||||
|
||||
|
||||
def _base_url(owned: OwnedProxy) -> str:
|
||||
return str(owned.gateway.client.base_url).rstrip("/")
|
||||
|
||||
|
||||
def _proxy_port(owned: OwnedProxy) -> int:
|
||||
return int(httpx.URL(_base_url(owned)).port or 0)
|
||||
|
||||
|
||||
def _accepted(pid: int, proxy_port: int, client_port: int) -> bool:
|
||||
return any(
|
||||
connection.status == psutil.CONN_ESTABLISHED
|
||||
and connection.laddr
|
||||
and connection.laddr.port == proxy_port
|
||||
and connection.raddr
|
||||
and connection.raddr.port == client_port
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
)
|
||||
|
||||
|
||||
def _listing_sample(owned: OwnedProxy, workers: tuple[int, ...]) -> tuple[int, tuple[str, ...]] | None:
|
||||
with httpx.Client(base_url=_base_url(owned), timeout=30, trust_env=False) as client:
|
||||
response: Final = client.get("/v1/models", headers={"Authorization": f"Bearer {owned.gateway.key}"})
|
||||
assert response.status_code == 200, response.text
|
||||
stream: Final[object] = response.extensions["network_stream"]
|
||||
assert isinstance(stream, httpcore.NetworkStream), stream
|
||||
client_port: Final = _ADDRESS.validate_python(stream.get_extra_info("client_addr"))[1]
|
||||
serving: Final = tuple(pid for pid in workers if _accepted(pid, _proxy_port(owned), client_port))
|
||||
return (serving[0], _model_ids(response.text)) if len(serving) == 1 else None
|
||||
|
||||
|
||||
def _listing_by_worker(owned: OwnedProxy, workers: tuple[int, ...]) -> Mapping[int, tuple[str, ...]]:
|
||||
samples: Final = tuple(_listing_sample(owned, workers) for _ in range(8))
|
||||
return MappingProxyType(dict(sample for sample in samples if sample is not None))
|
||||
|
||||
|
||||
def _login_names_listed(listed: tuple[str, ...]) -> frozenset[str]:
|
||||
return frozenset(name for name in listed if name in _FIXED or name.startswith(_LOGIN_PREFIXES))
|
||||
|
||||
|
||||
def _lists_login_models(listed: tuple[str, ...]) -> bool:
|
||||
login_listed: Final = _login_names_listed(listed)
|
||||
return frozenset(_FIXED) <= login_listed and all(
|
||||
pattern in login_listed or example in login_listed for pattern, example in _WILDCARDS.items()
|
||||
)
|
||||
|
||||
|
||||
def _every_worker_lists(
|
||||
owned: OwnedProxy, workers: tuple[int, ...], *, login_models: bool
|
||||
) -> Mapping[int, tuple[str, ...]]:
|
||||
def settled(by_worker: Mapping[int, tuple[str, ...]]) -> bool:
|
||||
if set(by_worker) != set(workers):
|
||||
return False
|
||||
if login_models:
|
||||
return all(_lists_login_models(listed) for listed in by_worker.values())
|
||||
return all(_login_names_listed(listed) == frozenset() for listed in by_worker.values())
|
||||
|
||||
return eventually(lambda: _listing_by_worker(owned, workers), settled, seconds=90)
|
||||
|
||||
|
||||
def _successes_logged(key: str, at_least: int) -> None:
|
||||
eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s AND status = %s',
|
||||
(hashlib.sha256(key.encode()).hexdigest(), "success"),
|
||||
),
|
||||
lambda rows: len(rows) >= at_least,
|
||||
seconds=40,
|
||||
)
|
||||
|
||||
|
||||
def _logged(served: _Served, status: str) -> None:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (served.call_id,)),
|
||||
lambda found: len(found) >= 1,
|
||||
seconds=40,
|
||||
)
|
||||
assert [str(row["status"]) for row in rows] == [status], rows
|
||||
|
||||
|
||||
def _reload(owned: OwnedProxy) -> None:
|
||||
response: Final = owned.gateway.request("POST", "/reload/model_cost_map", None)
|
||||
assert response.status_code == 200, response.text
|
||||
|
||||
|
||||
def _refusal_before_the_reload(name: str, universal: bool) -> str:
|
||||
if name in _FIXED:
|
||||
return _NO_HEALTHY
|
||||
if not universal:
|
||||
return _NOT_FOUND
|
||||
return dl.CHATGPT_REFUSAL if name.startswith("chatgpt/") else dl.COPILOT_REFUSAL
|
||||
|
||||
|
||||
@pytest.mark.parametrize("universal", _SHAPES.values(), ids=_SHAPES.keys())
|
||||
def test_booting_without_tokens_drops_the_login_deployments_and_a_reload_after_mounting_brings_them_back(
|
||||
gateway: Gateway, tmp_path: Path, universal: bool
|
||||
) -> None:
|
||||
dirs: Final = _token_dirs(tmp_path)
|
||||
with dl.device_login_peers(tmp_path) as peers:
|
||||
peers.switches.grant.set()
|
||||
with owned_proxy_process(
|
||||
gateway,
|
||||
tmp_path,
|
||||
{
|
||||
**peers.environment(),
|
||||
**dirs.environment(),
|
||||
"REQUEST_TIMEOUT": str(_REQUEST_TIMEOUT_SECONDS),
|
||||
"PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": str(_POLL_SECONDS),
|
||||
},
|
||||
config=_config(peers.api.url, tmp_path, fixed=True, wildcards=True, universal=universal),
|
||||
workers=2,
|
||||
) as owned:
|
||||
workers: Final = _workers(owned)
|
||||
boot_log: Final = owned.log.read_text()
|
||||
assert peers.auth_connections() == ()
|
||||
assert dl.CHATGPT_USER_CODE not in boot_log and dl.COPILOT_USER_CODE not in boot_log, boot_log[-3000:]
|
||||
assert sorted(_DROPPED.findall(boot_log)) == sorted([*_FIXED.values(), *_WILDCARDS] * 2), boot_log[-3000:]
|
||||
_every_worker_lists(owned, workers, login_models=False)
|
||||
with owned.gateway.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
for name in _LOGIN_NAMES:
|
||||
refused: Final = _chat(_base_url(owned), key, name, uuid.uuid4().hex)
|
||||
assert refused.status == 400, refused.text
|
||||
assert _refusal_before_the_reload(name, universal) in refused.message(), refused.text
|
||||
control: Final = _chat(_base_url(owned), key, _CONTROL, uuid.uuid4().hex)
|
||||
assert control.status == 200, control.text
|
||||
assert peers.auth_connections() == ()
|
||||
assert dl.bearer(peers.api.drain()[-1]) == dl.CONTROL_KEY
|
||||
|
||||
dirs.store_valid_tokens(peers.api.url)
|
||||
_reload(owned)
|
||||
_every_worker_lists(owned, workers, login_models=True)
|
||||
for served_logins, name in enumerate(_LOGIN_NAMES, start=1):
|
||||
served: Final = _chat(_base_url(owned), key, name, uuid.uuid4().hex)
|
||||
assert served.status == 200, served.text
|
||||
_successes_logged(key, at_least=_CONTROL_CALLS_BEFORE_THE_RELOAD + served_logins)
|
||||
assert peers.auth_connections() == ()
|
||||
forwarded: Final = peers.api.drain()
|
||||
assert {dl.bearer(request) for request in forwarded} == {dl.CHATGPT_STORED, dl.COPILOT_STORED_KEY}
|
||||
|
||||
dl.write_chatgpt(dirs.chatgpt, dl.chatgpt_record(dl.CHATGPT_EXPIRED, refresh_token=dl.GOOD_REFRESH))
|
||||
peers.switches.hold_connect.set()
|
||||
with ThreadPoolExecutor(max_workers=1) as pool:
|
||||
pending: Final = pool.submit(_chat, _base_url(owned), key, "chatgpt-fixed", uuid.uuid4().hex, 120)
|
||||
try:
|
||||
dropped: Final = eventually(peers.dropped, lambda hangups: len(hangups) >= 1, seconds=40)
|
||||
finally:
|
||||
peers.switches.refuse_connect.set()
|
||||
peers.switches.hold_connect.clear()
|
||||
timed_out: Final = pending.result(timeout=120)
|
||||
assert dropped[0].authority == f"{dl.CHATGPT_AUTH_HOST}:443", dropped
|
||||
assert _REQUEST_TIMEOUT_SECONDS - 2 <= dropped[0].seconds <= _REQUEST_TIMEOUT_SECONDS + 6, dropped
|
||||
assert timed_out.status == 400, timed_out.text
|
||||
assert dl.CHATGPT_REFUSAL in timed_out.message(), timed_out.text
|
||||
_logged(timed_out, "failure")
|
||||
peers.switches.refuse_connect.clear()
|
||||
assert set(peers.auth_connections()) <= {f"{dl.CHATGPT_AUTH_HOST}:443"}
|
||||
assert peers.dropped() == ()
|
||||
|
||||
marker: Final = uuid.uuid4().hex
|
||||
recovered: Final = _chat(_base_url(owned), key, "chatgpt-fixed", marker)
|
||||
assert recovered.status == 200, recovered.text
|
||||
assert set(rv.MARKER.findall(recovered.text)) == {marker}, recovered.text
|
||||
refreshes: Final = [(request.method, request.target) for request in peers.auth.drain()]
|
||||
assert refreshes == [("POST", "/oauth/token")], refreshes
|
||||
assert dl.bearer(peers.api.drain()[-1]) == dl.CHATGPT_REFRESHED
|
||||
|
||||
|
||||
@pytest.mark.parametrize("shape", ("fixed", "wildcards"))
|
||||
def test_a_token_removed_after_boot_turns_every_request_into_the_refusal(
|
||||
gateway: Gateway, tmp_path: Path, shape: str
|
||||
) -> None:
|
||||
dirs: Final = _token_dirs(tmp_path)
|
||||
with dl.device_login_peers(tmp_path) as peers:
|
||||
dirs.store_valid_tokens(peers.api.url)
|
||||
with owned_proxy_process(
|
||||
gateway,
|
||||
tmp_path,
|
||||
{**peers.environment(), **dirs.environment()},
|
||||
config=_config(peers.api.url, tmp_path, fixed=shape == "fixed", wildcards=shape == "wildcards"),
|
||||
workers=2,
|
||||
) as owned:
|
||||
_workers(owned)
|
||||
names: Final = tuple(_FIXED) if shape == "fixed" else tuple(_FIXED.values())
|
||||
with owned.gateway.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
for name in names:
|
||||
served: Final = _chat(_base_url(owned), key, name, uuid.uuid4().hex)
|
||||
assert served.status == 200, served.text
|
||||
dl.clear_tokens(dirs.chatgpt, dirs.copilot)
|
||||
peers.switches.grant.set()
|
||||
for name in names:
|
||||
refusal: Final = dl.CHATGPT_REFUSAL if name.startswith("chatgpt") else dl.COPILOT_REFUSAL
|
||||
for _ in range(4):
|
||||
refused: Final = _chat(_base_url(owned), key, name, uuid.uuid4().hex)
|
||||
assert refused.status == 400, refused.text
|
||||
assert refusal in refused.message(), refused.text
|
||||
assert peers.auth_connections() == ()
|
||||
tail: Final = owned.log.read_text()
|
||||
assert dl.CHATGPT_USER_CODE not in tail and dl.COPILOT_USER_CODE not in tail, tail[-3000:]
|
||||
control: Final = _chat(_base_url(owned), key, _CONTROL, uuid.uuid4().hex)
|
||||
assert control.status == 200, control.text
|
||||
|
||||
|
||||
async def _send(client: httpx.AsyncClient, key: str, model: str, marker: str) -> _Served | None:
|
||||
try:
|
||||
response: Final = await client.post(
|
||||
"/v1/chat/completions",
|
||||
json={"model": model, "messages": [{"role": "user", "content": f"Reply to marker-{marker}"}], **_EXTRA},
|
||||
headers={"Authorization": f"Bearer {key}"},
|
||||
)
|
||||
except httpx.TransportError:
|
||||
return None
|
||||
return _Served(response.status_code, response.text, response.headers.get(_CALL_ID, ""))
|
||||
|
||||
|
||||
async def _burst(base_url: str, key: str, models: tuple[str, ...]) -> tuple[_Served | None, ...]:
|
||||
async with httpx.AsyncClient(base_url=base_url, timeout=120, trust_env=False) as client:
|
||||
return tuple(await asyncio.gather(*(_send(client, key, model, uuid.uuid4().hex) for model in models)))
|
||||
|
||||
|
||||
def _open_upstream_connections(pid: int, port: int) -> int:
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
||||
|
||||
_Bursts: TypeAlias = tuple[asyncio.Task[tuple[_Served | None, ...]], ...]
|
||||
_HELD_PER_BURST: Final = 8
|
||||
_HELD_PER_WORKER: Final = 4
|
||||
_MOST_BURSTS: Final = 8
|
||||
|
||||
|
||||
async def _held_on_every_worker(
|
||||
owned: OwnedProxy,
|
||||
key: str,
|
||||
workers: tuple[int, ...],
|
||||
control_port: int,
|
||||
held_markers: SimpleQueue[str],
|
||||
bursts: _Bursts,
|
||||
) -> tuple[_Bursts, Mapping[int, int]]:
|
||||
started: Final = (*bursts, asyncio.create_task(_burst(_base_url(owned), key, (_CONTROL,) * _HELD_PER_BURST)))
|
||||
expected: Final = _HELD_PER_BURST * len(started)
|
||||
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == expected, 60)
|
||||
held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, control_port) for pid in workers})
|
||||
assert sum(held_by.values()) == expected, held_by
|
||||
if min(held_by.values()) >= _HELD_PER_WORKER:
|
||||
return started, held_by
|
||||
assert len(started) < _MOST_BURSTS, held_by
|
||||
return await _held_on_every_worker(owned, key, workers, control_port, held_markers, started)
|
||||
|
||||
|
||||
async def test_worker_sigkill_mid_burst_leaves_the_sibling_refusing_logins_and_serving_the_control(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
dirs: Final = _token_dirs(tmp_path)
|
||||
release: Final = threading.Event()
|
||||
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
vendor: Final = rv.ResponsesVendor()
|
||||
|
||||
def held(request: Request) -> Reply:
|
||||
if request.method == "GET":
|
||||
return vendor.respond(request)
|
||||
held_markers.put(rv.newest_marker(request.body.decode()) or "")
|
||||
assert release.wait(timeout=120), "The burst was never released"
|
||||
return vendor.respond(request)
|
||||
|
||||
with (
|
||||
dl.device_login_peers(tmp_path) as peers,
|
||||
wire_server(held) as control_api,
|
||||
owned_proxy_process(
|
||||
gateway,
|
||||
tmp_path,
|
||||
{**peers.environment(), **dirs.environment()},
|
||||
config=_config(control_api.url, tmp_path, fixed=False, wildcards=False, universal=True),
|
||||
workers=2,
|
||||
) as owned,
|
||||
):
|
||||
workers: Final = _workers(owned)
|
||||
control_port: Final = int(control_api.url.rsplit(":", 1)[1])
|
||||
with owned.gateway.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
bursts, held_by = await _held_on_every_worker(owned, key, workers, control_port, held_markers, ())
|
||||
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
|
||||
victim: Final = psutil.Process(victim_pid)
|
||||
victim.suspend()
|
||||
victim.send_signal(signal.SIGKILL)
|
||||
release.set()
|
||||
answers: Final = await asyncio.gather(*bursts)
|
||||
served: Final = tuple(item for item in itertools.chain.from_iterable(answers) if item is not None)
|
||||
assert len(served) == held_by[survivor_pid], (held_by, len(served))
|
||||
for item in served:
|
||||
assert item.status == 200, item.text
|
||||
refusals: Final = await _burst(
|
||||
_base_url(owned), key, ("chatgpt/gpt-5.6-terra", "github_copilot/gpt-5.2") * 10
|
||||
)
|
||||
for index, answer in enumerate(refusals):
|
||||
assert answer is not None and answer.status == 400, answer
|
||||
refusal: Final = dl.CHATGPT_REFUSAL if index % 2 == 0 else dl.COPILOT_REFUSAL
|
||||
assert refusal in answer.message(), answer.text
|
||||
assert peers.auth_connections() == ()
|
||||
follow_up: Final = _chat(_base_url(owned), key, _CONTROL, uuid.uuid4().hex)
|
||||
assert follow_up.status == 200, follow_up.text
|
||||
899
tests/integration/providers/test_device_code_login_guard_wire.py
Normal file
899
tests/integration/providers/test_device_code_login_guard_wire.py
Normal file
|
|
@ -0,0 +1,899 @@
|
|||
import asyncio
|
||||
import hashlib
|
||||
import itertools
|
||||
import json
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support import device_login as dl
|
||||
from integration._support import responses_vendor as rv
|
||||
from integration._support.client import eventually, gateway_from_environment
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import OwnedProxy, owned_proxy_process
|
||||
from integration._support.wire import Request
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(300)
|
||||
|
||||
Provider: TypeAlias = Literal["chatgpt", "copilot"]
|
||||
Endpoint: TypeAlias = Literal["chat", "responses", "messages"]
|
||||
Client: TypeAlias = Literal["sdk", "async_sdk", "httpx"]
|
||||
|
||||
_PROVIDERS: Final[tuple[Provider, ...]] = ("chatgpt", "copilot")
|
||||
_ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "responses", "messages")
|
||||
_CLIENTS: Final[tuple[Client, ...]] = ("sdk", "async_sdk", "httpx")
|
||||
_MODELS: Final[Mapping[tuple[Provider, Endpoint], str]] = MappingProxyType(
|
||||
{
|
||||
("chatgpt", "chat"): "chatgpt/gpt-5.6-terra",
|
||||
("chatgpt", "responses"): "chatgpt/gpt-5.6-terra",
|
||||
("chatgpt", "messages"): "chatgpt/gpt-5.6-terra",
|
||||
("copilot", "chat"): "github_copilot/gpt-5.2",
|
||||
("copilot", "responses"): "github_copilot/gpt-5.2",
|
||||
("copilot", "messages"): "github_copilot/claude-sonnet-4.5",
|
||||
}
|
||||
)
|
||||
_REFUSALS: Final[Mapping[Provider, str]] = MappingProxyType(
|
||||
{"chatgpt": dl.CHATGPT_REFUSAL, "copilot": dl.COPILOT_REFUSAL}
|
||||
)
|
||||
_STORED_BEARERS: Final[Mapping[Provider, str]] = MappingProxyType(
|
||||
{"chatgpt": dl.CHATGPT_STORED, "copilot": dl.COPILOT_STORED_KEY}
|
||||
)
|
||||
_CONTROL_MODEL: Final = "device-login-control"
|
||||
_EXTRA: Final[Mapping[str, JsonValue]] = MappingProxyType({"num_retries": 0, "cache": {"no-cache": True}})
|
||||
_CALL_ID: Final = "x-litellm-call-id"
|
||||
_CLIENT_SECONDS: Final = 60
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Cell:
|
||||
provider: Provider
|
||||
endpoint: Endpoint
|
||||
stream: bool
|
||||
client: Client
|
||||
|
||||
def name(self) -> str:
|
||||
return f"{self.provider}-{self.endpoint}-{'stream' if self.stream else 'unary'}-{self.client}"
|
||||
|
||||
|
||||
_CELLS: Final = tuple(
|
||||
_Cell(provider, endpoint, stream, client)
|
||||
for provider, endpoint, stream, client in itertools.product(_PROVIDERS, _ENDPOINTS, (False, True), _CLIENTS)
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Call:
|
||||
base_url: str
|
||||
key: str
|
||||
model: str
|
||||
endpoint: Endpoint
|
||||
stream: bool
|
||||
client: Client
|
||||
marker: str
|
||||
served: SimpleQueue[str]
|
||||
|
||||
def prompt(self) -> str:
|
||||
return f"Reply to marker-{self.marker}"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
status: int
|
||||
text: str
|
||||
call_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Rig:
|
||||
owned: OwnedProxy
|
||||
peers: dl.Peers
|
||||
chatgpt_dir: Path
|
||||
copilot_dir: Path
|
||||
|
||||
def base_url(self) -> str:
|
||||
return str(self.owned.gateway.client.base_url).rstrip("/")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Scene:
|
||||
rig: _Rig
|
||||
key: str
|
||||
log_offset: int
|
||||
served: SimpleQueue[str]
|
||||
|
||||
def call(self, cell: _Cell) -> _Call:
|
||||
return _Call(
|
||||
self.rig.base_url(),
|
||||
self.key,
|
||||
_MODELS[cell.provider, cell.endpoint],
|
||||
cell.endpoint,
|
||||
cell.stream,
|
||||
cell.client,
|
||||
uuid.uuid4().hex,
|
||||
self.served,
|
||||
)
|
||||
|
||||
def control(self, stream: bool) -> _Call:
|
||||
return _Call(
|
||||
self.rig.base_url(), self.key, _CONTROL_MODEL, "chat", stream, "httpx", uuid.uuid4().hex, self.served
|
||||
)
|
||||
|
||||
def post(self, path: str, body: Mapping[str, JsonValue], *, as_admin: bool = False) -> _Served:
|
||||
key: Final = self.rig.owned.gateway.key if as_admin else self.key
|
||||
with httpx.Client(base_url=self.rig.base_url(), timeout=_CLIENT_SECONDS, trust_env=False) as client:
|
||||
response: Final = client.post(
|
||||
path, json=dict(body), headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"}
|
||||
)
|
||||
return _Served(response.status_code, response.text, response.headers.get(_CALL_ID, ""))
|
||||
|
||||
def token_files(self) -> tuple[str, ...]:
|
||||
return tuple(sorted(path.name for path in (*self.rig.chatgpt_dir.iterdir(), *self.rig.copilot_dir.iterdir())))
|
||||
|
||||
def log_tail(self) -> str:
|
||||
with self.rig.owned.log.open("rb") as log:
|
||||
log.seek(self.log_offset)
|
||||
return log.read().decode(errors="replace")
|
||||
|
||||
def store_valid_token(self, provider: Provider) -> None:
|
||||
if provider == "chatgpt":
|
||||
dl.write_chatgpt(self.rig.chatgpt_dir, dl.chatgpt_record(dl.CHATGPT_STORED))
|
||||
return
|
||||
dl.write_copilot_key(self.rig.copilot_dir, dl.COPILOT_STORED_KEY, self.rig.peers.api.url)
|
||||
|
||||
|
||||
def _config(peers: dl.Peers, directory: Path) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{"model_name": "*", "litellm_params": {"model": "*"}},
|
||||
{
|
||||
"model_name": _CONTROL_MODEL,
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": f"{peers.api.url}/v1",
|
||||
"api_key": dl.CONTROL_KEY,
|
||||
},
|
||||
},
|
||||
]
|
||||
path: Final = directory / "device-login-guard.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]:
|
||||
directory: Final = tmp_path_factory.mktemp("device-login-guard")
|
||||
chatgpt_dir: Final = directory / "chatgpt"
|
||||
copilot_dir: Final = directory / "copilot"
|
||||
chatgpt_dir.mkdir()
|
||||
copilot_dir.mkdir()
|
||||
with (
|
||||
gateway_from_environment() as gateway,
|
||||
dl.device_login_peers(directory) as peers,
|
||||
owned_proxy_process(
|
||||
gateway,
|
||||
directory,
|
||||
{
|
||||
**peers.environment(),
|
||||
"CHATGPT_TOKEN_DIR": str(chatgpt_dir),
|
||||
"GITHUB_COPILOT_TOKEN_DIR": str(copilot_dir),
|
||||
},
|
||||
config=_config(peers, directory),
|
||||
workers=2,
|
||||
remove_environment=("REQUEST_TIMEOUT",),
|
||||
) as owned,
|
||||
):
|
||||
yield _Rig(owned, peers, chatgpt_dir, copilot_dir)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def scene(rig: _Rig) -> Iterator[_Scene]:
|
||||
rig.peers.reset()
|
||||
dl.clear_tokens(rig.chatgpt_dir, rig.copilot_dir)
|
||||
served: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
with rig.owned.gateway.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
yield _Scene(rig, key, rig.owned.log.stat().st_size, served)
|
||||
_every_served_call_logged(key, served)
|
||||
rig.peers.reset()
|
||||
dl.clear_tokens(rig.chatgpt_dir, rig.copilot_dir)
|
||||
|
||||
|
||||
def _path(endpoint: Endpoint) -> str:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return "/v1/chat/completions"
|
||||
case "responses":
|
||||
return "/v1/responses"
|
||||
case "messages":
|
||||
return "/v1/messages"
|
||||
|
||||
|
||||
def _raw_body(call: _Call) -> Mapping[str, JsonValue]:
|
||||
common: Final[Mapping[str, JsonValue]] = MappingProxyType({"model": call.model, "stream": call.stream, **_EXTRA})
|
||||
match call.endpoint:
|
||||
case "chat":
|
||||
return {**common, "messages": [{"role": "user", "content": call.prompt()}]}
|
||||
case "responses":
|
||||
return {**common, "input": call.prompt()}
|
||||
case "messages":
|
||||
return {**common, "max_tokens": 64, "messages": [{"role": "user", "content": call.prompt()}]}
|
||||
|
||||
|
||||
def _httpx(call: _Call, seconds: float = _CLIENT_SECONDS) -> _Served:
|
||||
with (
|
||||
httpx.Client(base_url=call.base_url, timeout=seconds, trust_env=False) as client,
|
||||
client.stream(
|
||||
"POST",
|
||||
_path(call.endpoint),
|
||||
json=_raw_body(call),
|
||||
headers={"Authorization": f"Bearer {call.key}", "anthropic-version": "2023-06-01"},
|
||||
) as response,
|
||||
):
|
||||
return _record(
|
||||
call, _Served(response.status_code, response.read().decode(), response.headers.get(_CALL_ID, ""))
|
||||
)
|
||||
|
||||
|
||||
def _refused(error: openai.APIStatusError | anthropic.APIStatusError) -> _Served:
|
||||
return _Served(error.status_code, error.response.text, error.response.headers.get(_CALL_ID, ""))
|
||||
|
||||
|
||||
def _sdk(call: _Call) -> _Served:
|
||||
try:
|
||||
if call.endpoint == "messages":
|
||||
with (
|
||||
anthropic.Anthropic(
|
||||
base_url=call.base_url, api_key=call.key, max_retries=0, timeout=_CLIENT_SECONDS
|
||||
) as claude,
|
||||
claude.messages.with_streaming_response.create(
|
||||
model=call.model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": call.prompt()}],
|
||||
stream=call.stream,
|
||||
extra_body=dict(_EXTRA),
|
||||
) as message,
|
||||
):
|
||||
return _Served(message.status_code, message.text(), message.headers.get(_CALL_ID, ""))
|
||||
with openai.OpenAI(
|
||||
base_url=f"{call.base_url}/v1", api_key=call.key, max_retries=0, timeout=_CLIENT_SECONDS
|
||||
) as client:
|
||||
if call.endpoint == "chat":
|
||||
with client.chat.completions.with_streaming_response.create(
|
||||
model=call.model,
|
||||
messages=[{"role": "user", "content": call.prompt()}],
|
||||
stream=call.stream,
|
||||
extra_body=dict(_EXTRA),
|
||||
) as completion:
|
||||
return _Served(completion.status_code, completion.text(), completion.headers.get(_CALL_ID, ""))
|
||||
with client.responses.with_streaming_response.create(
|
||||
model=call.model, input=call.prompt(), stream=call.stream, extra_body=dict(_EXTRA)
|
||||
) as created:
|
||||
return _Served(created.status_code, created.text(), created.headers.get(_CALL_ID, ""))
|
||||
except (openai.APIStatusError, anthropic.APIStatusError) as error:
|
||||
return _refused(error)
|
||||
|
||||
|
||||
async def _async_sdk(call: _Call) -> _Served:
|
||||
try:
|
||||
if call.endpoint == "messages":
|
||||
async with (
|
||||
anthropic.AsyncAnthropic(
|
||||
base_url=call.base_url, api_key=call.key, max_retries=0, timeout=_CLIENT_SECONDS
|
||||
) as claude,
|
||||
claude.messages.with_streaming_response.create(
|
||||
model=call.model,
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": call.prompt()}],
|
||||
stream=call.stream,
|
||||
extra_body=dict(_EXTRA),
|
||||
) as message,
|
||||
):
|
||||
return _Served(message.status_code, await message.text(), message.headers.get(_CALL_ID, ""))
|
||||
async with openai.AsyncOpenAI(
|
||||
base_url=f"{call.base_url}/v1", api_key=call.key, max_retries=0, timeout=_CLIENT_SECONDS
|
||||
) as client:
|
||||
if call.endpoint == "chat":
|
||||
async with client.chat.completions.with_streaming_response.create(
|
||||
model=call.model,
|
||||
messages=[{"role": "user", "content": call.prompt()}],
|
||||
stream=call.stream,
|
||||
extra_body=dict(_EXTRA),
|
||||
) as completion:
|
||||
return _Served(
|
||||
completion.status_code, await completion.text(), completion.headers.get(_CALL_ID, "")
|
||||
)
|
||||
async with client.responses.with_streaming_response.create(
|
||||
model=call.model, input=call.prompt(), stream=call.stream, extra_body=dict(_EXTRA)
|
||||
) as created:
|
||||
return _Served(created.status_code, await created.text(), created.headers.get(_CALL_ID, ""))
|
||||
except (openai.APIStatusError, anthropic.APIStatusError) as error:
|
||||
return _refused(error)
|
||||
|
||||
|
||||
def _serve(call: _Call) -> _Served:
|
||||
match call.client:
|
||||
case "sdk":
|
||||
return _record(call, _sdk(call))
|
||||
case "async_sdk":
|
||||
return _record(call, asyncio.run(_async_sdk(call)))
|
||||
case "httpx":
|
||||
return _httpx(call)
|
||||
|
||||
|
||||
def _error_message(text: str) -> str:
|
||||
error: Final = rv.JSON_OBJECT.validate_python(rv.JSON_OBJECT.validate_json(text)["error"])
|
||||
return str(error["message"])
|
||||
|
||||
|
||||
def _spend_rows(key: str) -> Sequence[Mapping[str, JsonValue]]:
|
||||
return eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, status, model FROM "LiteLLM_SpendLogs" WHERE api_key = %s',
|
||||
(hashlib.sha256(key.encode()).hexdigest(),),
|
||||
),
|
||||
lambda rows: len(rows) >= 1,
|
||||
seconds=40,
|
||||
)
|
||||
|
||||
|
||||
def _record(call: _Call, served: _Served) -> _Served:
|
||||
call.served.put(served.call_id)
|
||||
return served
|
||||
|
||||
|
||||
def _every_served_call_logged(key: str, served: SimpleQueue[str]) -> None:
|
||||
served_calls: Final = tuple(served.get_nowait() for _ in range(served.qsize()))
|
||||
if not served_calls:
|
||||
return
|
||||
eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s',
|
||||
(hashlib.sha256(key.encode()).hexdigest(),),
|
||||
),
|
||||
lambda rows: len(rows) >= len(served_calls),
|
||||
seconds=40,
|
||||
)
|
||||
|
||||
|
||||
def _failure_logged(served: _Served) -> None:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (served.call_id,)),
|
||||
lambda found: len(found) >= 1,
|
||||
seconds=40,
|
||||
)
|
||||
assert [str(row["status"]) for row in rows] == ["failure"], rows
|
||||
|
||||
|
||||
def _carrying(received: tuple[Request, ...], marker: str) -> tuple[Request, ...]:
|
||||
return tuple(request for request in received if marker.encode() in request.body)
|
||||
|
||||
|
||||
def _describe(received: tuple[Request, ...]) -> tuple[tuple[str, str], ...]:
|
||||
return tuple((request.method, request.target) for request in received)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("cell", _CELLS, ids=_Cell.name)
|
||||
def test_missing_token_is_refused_with_the_login_instructions_before_any_auth_host_call(
|
||||
scene: _Scene, cell: _Cell
|
||||
) -> None:
|
||||
call: Final = scene.call(cell)
|
||||
served: Final = _serve(call)
|
||||
assert served.status == 400, served.text
|
||||
assert _REFUSALS[cell.provider] in _error_message(served.text), served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
(row,) = _spend_rows(scene.key)
|
||||
assert (row["request_id"], row["status"]) == (served.call_id, "failure"), row
|
||||
|
||||
|
||||
@pytest.mark.parametrize("cell", _CELLS, ids=_Cell.name)
|
||||
def test_stored_token_serves_the_request_without_an_auth_host_call(scene: _Scene, cell: _Cell) -> None:
|
||||
scene.store_valid_token(cell.provider)
|
||||
call: Final = scene.call(cell)
|
||||
served: Final = _serve(call)
|
||||
assert served.status == 200, served.text
|
||||
assert set(rv.MARKER.findall(served.text)) == {call.marker}, served.text
|
||||
received: Final = scene.rig.peers.api.drain()
|
||||
(forwarded,) = _carrying(received, call.marker)
|
||||
assert dl.bearer(forwarded) == _STORED_BEARERS[cell.provider], _describe(received)
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
(row,) = _spend_rows(scene.key)
|
||||
assert row["status"] == "success", row
|
||||
|
||||
|
||||
def test_expired_chatgpt_token_is_refreshed_once_and_the_request_is_served(scene: _Scene) -> None:
|
||||
dl.write_chatgpt(scene.rig.chatgpt_dir, dl.chatgpt_record(dl.CHATGPT_EXPIRED, refresh_token=dl.GOOD_REFRESH))
|
||||
call: Final = scene.call(_Cell("chatgpt", "responses", False, "httpx"))
|
||||
served: Final = _serve(call)
|
||||
assert served.status == 200, served.text
|
||||
assert set(rv.MARKER.findall(served.text)) == {call.marker}, served.text
|
||||
(refresh,) = scene.rig.peers.auth.drain()
|
||||
assert (refresh.method, refresh.target, refresh.headers.get("host")) == (
|
||||
"POST",
|
||||
"/oauth/token",
|
||||
dl.CHATGPT_AUTH_HOST,
|
||||
)
|
||||
sent: Final = rv.JSON_OBJECT.validate_json(refresh.body)
|
||||
assert (sent["grant_type"], sent["refresh_token"]) == ("refresh_token", dl.GOOD_REFRESH), sent
|
||||
(forwarded,) = _carrying(scene.rig.peers.api.drain(), call.marker)
|
||||
assert dl.bearer(forwarded) == dl.CHATGPT_REFRESHED
|
||||
stored: Final = json.loads((scene.rig.chatgpt_dir / "auth.json").read_text())
|
||||
assert stored["access_token"] == dl.CHATGPT_REFRESHED, sorted(stored)
|
||||
|
||||
|
||||
def test_revoked_chatgpt_refresh_token_is_refused_after_the_refresh_call_alone(scene: _Scene) -> None:
|
||||
dl.write_chatgpt(scene.rig.chatgpt_dir, dl.chatgpt_record(dl.CHATGPT_EXPIRED, refresh_token=dl.REVOKED_REFRESH))
|
||||
served: Final = _serve(scene.call(_Cell("chatgpt", "chat", False, "httpx")))
|
||||
assert served.status == 400, served.text
|
||||
assert dl.CHATGPT_REFUSAL in _error_message(served.text), served.text
|
||||
_failure_logged(served)
|
||||
reached: Final = scene.rig.peers.auth.drain()
|
||||
assert {(request.method, request.target) for request in reached} == {("POST", "/oauth/token")}, _describe(reached)
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Route:
|
||||
name: str
|
||||
provider: Provider
|
||||
path: str
|
||||
body: Mapping[str, JsonValue]
|
||||
|
||||
def label(self) -> str:
|
||||
return self.name
|
||||
|
||||
|
||||
_ROUTES: Final = (
|
||||
_Route("chatgpt-embeddings", "chatgpt", "/v1/embeddings", {"model": "chatgpt/gpt-5.6-terra", "input": "login"}),
|
||||
_Route("chatgpt-completions", "chatgpt", "/v1/completions", {"model": "chatgpt/gpt-5.6-terra", "prompt": "login"}),
|
||||
_Route(
|
||||
"chatgpt-images", "chatgpt", "/v1/images/generations", {"model": "chatgpt/gpt-5.6-terra", "prompt": "login"}
|
||||
),
|
||||
_Route(
|
||||
"copilot-embeddings",
|
||||
"copilot",
|
||||
"/v1/embeddings",
|
||||
{"model": "github_copilot/text-embedding-3-small", "input": "login"},
|
||||
),
|
||||
_Route("copilot-completions", "copilot", "/v1/completions", {"model": "github_copilot/gpt-5.2", "prompt": "login"}),
|
||||
_Route(
|
||||
"copilot-images", "copilot", "/v1/images/generations", {"model": "github_copilot/gpt-5.2", "prompt": "login"}
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", _ROUTES, ids=_Route.label)
|
||||
def test_missing_token_is_refused_on_the_other_model_routes(scene: _Scene, route: _Route) -> None:
|
||||
served: Final = scene.post(route.path, {**route.body, **_EXTRA})
|
||||
assert served.status == 400, served.text
|
||||
assert _REFUSALS[route.provider] in _error_message(served.text), served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
|
||||
|
||||
def test_copilot_embeddings_are_served_with_the_stored_api_key(scene: _Scene) -> None:
|
||||
scene.store_valid_token("copilot")
|
||||
served: Final = scene.post(
|
||||
"/v1/embeddings", {"model": "github_copilot/text-embedding-3-small", "input": "login", **_EXTRA}
|
||||
)
|
||||
assert served.status == 200, served.text
|
||||
(item,) = rv.ITEMS.validate_python(rv.JSON_OBJECT.validate_json(served.text)["data"])
|
||||
assert item["embedding"] == [0.25, 0.5, 0.75], served.text
|
||||
(forwarded,) = scene.rig.peers.api.drain()
|
||||
assert (forwarded.method, dl.bearer(forwarded)) == ("POST", dl.COPILOT_STORED_KEY), forwarded.target
|
||||
assert forwarded.target.endswith("/embeddings"), forwarded.target
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", _PROVIDERS)
|
||||
def test_connection_test_reports_the_refusal_instead_of_starting_a_login(scene: _Scene, provider: Provider) -> None:
|
||||
served: Final = scene.post(
|
||||
"/health/test_connection",
|
||||
{"litellm_params": {"model": _MODELS[provider, "chat"]}, "mode": "chat"},
|
||||
as_admin=True,
|
||||
)
|
||||
assert served.status == 200, served.text
|
||||
report: Final = rv.JSON_OBJECT.validate_json(served.text)
|
||||
assert report["status"] == "error", served.text
|
||||
assert _REFUSALS[provider] in str(rv.JSON_OBJECT.validate_python(report["result"])["error"]), served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", _PROVIDERS)
|
||||
def test_count_tokens_answers_locally_without_a_login(scene: _Scene, provider: Provider) -> None:
|
||||
served: Final = scene.post(
|
||||
"/v1/messages/count_tokens",
|
||||
{"model": _MODELS[provider, "messages"], "messages": [{"role": "user", "content": "count these tokens"}]},
|
||||
)
|
||||
assert served.status == 200, served.text
|
||||
counted: Final = rv.JSON_OBJECT.validate_json(served.text)["input_tokens"]
|
||||
assert isinstance(counted, int) and counted > 0, served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
|
||||
|
||||
def test_health_check_still_answers_without_a_login(scene: _Scene) -> None:
|
||||
response: Final = scene.rig.owned.gateway.request("GET", "/health", None)
|
||||
assert response.status_code == 200, response.text
|
||||
health: Final = rv.JSON_OBJECT.validate_json(response.text)
|
||||
healthy: Final = rv.ITEMS.validate_python(health["healthy_endpoints"])
|
||||
assert [endpoint["model"] for endpoint in healthy] == ["openai/gpt-4o-mini"], response.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
tail: Final = scene.log_tail()
|
||||
assert dl.CHATGPT_USER_CODE not in tail and dl.COPILOT_USER_CODE not in tail, tail[-2000:]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", _ENDPOINTS)
|
||||
@pytest.mark.parametrize("provider", _PROVIDERS)
|
||||
def test_an_auth_host_ready_to_grant_a_login_is_never_asked_and_no_code_reaches_the_log(
|
||||
scene: _Scene, provider: Provider, endpoint: Endpoint
|
||||
) -> None:
|
||||
scene.rig.peers.switches.grant.set()
|
||||
served: Final = _serve(scene.call(_Cell(provider, endpoint, False, "httpx")))
|
||||
assert served.status == 400, served.text
|
||||
assert _REFUSALS[provider] in _error_message(served.text), served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
assert _describe(scene.rig.peers.auth.drain()) == ()
|
||||
tail: Final = scene.log_tail()
|
||||
assert dl.CHATGPT_USER_CODE not in tail and dl.COPILOT_USER_CODE not in tail, tail[-2000:]
|
||||
assert scene.token_files() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", _ENDPOINTS)
|
||||
@pytest.mark.parametrize("provider", _PROVIDERS)
|
||||
def test_stored_token_the_provider_rejects_surfaces_its_401_without_a_login(
|
||||
scene: _Scene, provider: Provider, endpoint: Endpoint
|
||||
) -> None:
|
||||
if provider == "chatgpt":
|
||||
dl.write_chatgpt(scene.rig.chatgpt_dir, dl.chatgpt_record(dl.CHATGPT_REJECTED))
|
||||
else:
|
||||
dl.write_copilot_key(scene.rig.copilot_dir, dl.COPILOT_REJECTED_KEY, scene.rig.peers.api.url)
|
||||
call: Final = scene.call(_Cell(provider, endpoint, False, "httpx"))
|
||||
served: Final = _serve(call)
|
||||
assert served.status == 401, served.text
|
||||
assert dl.REJECTED_DETAIL in served.text, served.text
|
||||
received: Final = scene.rig.peers.api.drain()
|
||||
assert {dl.bearer(request) for request in _carrying(received, call.marker)} == {
|
||||
dl.CHATGPT_REJECTED if provider == "chatgpt" else dl.COPILOT_REJECTED_KEY
|
||||
}, _describe(received)
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stale_key", (False, True), ids=("missing-api-key", "expired-api-key"))
|
||||
def test_copilot_api_key_is_minted_once_from_the_stored_access_token(scene: _Scene, stale_key: bool) -> None:
|
||||
dl.write_copilot_access(scene.rig.copilot_dir, dl.COPILOT_ACCESS)
|
||||
if stale_key:
|
||||
dl.write_copilot_key(scene.rig.copilot_dir, dl.COPILOT_REJECTED_KEY, scene.rig.peers.api.url, -3600)
|
||||
call: Final = scene.call(_Cell("copilot", "chat", False, "httpx"))
|
||||
served: Final = _serve(call)
|
||||
assert served.status == 200, served.text
|
||||
assert set(rv.MARKER.findall(served.text)) == {call.marker}, served.text
|
||||
(minted,) = scene.rig.peers.auth.drain()
|
||||
assert (minted.method, minted.target, minted.headers.get("host"), minted.headers.get("authorization")) == (
|
||||
"GET",
|
||||
"/copilot_internal/v2/token",
|
||||
dl.GITHUB_API_HOST,
|
||||
f"token {dl.COPILOT_ACCESS}",
|
||||
)
|
||||
(forwarded,) = _carrying(scene.rig.peers.api.drain(), call.marker)
|
||||
assert dl.bearer(forwarded) == dl.COPILOT_MINTED_KEY
|
||||
assert json.loads((scene.rig.copilot_dir / "api-key.json").read_text())["token"] == dl.COPILOT_MINTED_KEY
|
||||
|
||||
|
||||
def test_copilot_access_token_github_rejects_fails_the_request_without_a_device_login(scene: _Scene) -> None:
|
||||
dl.write_copilot_access(scene.rig.copilot_dir, dl.COPILOT_REJECTED_ACCESS)
|
||||
served: Final = _serve(scene.call(_Cell("copilot", "chat", False, "httpx")))
|
||||
assert served.status == 400, served.text
|
||||
assert "Failed to refresh API key" in _error_message(served.text), served.text
|
||||
_failure_logged(served)
|
||||
reached: Final = scene.rig.peers.auth.drain()
|
||||
assert {(request.method, request.target) for request in reached} == {("GET", "/copilot_internal/v2/token")}, (
|
||||
_describe(reached)
|
||||
)
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
|
||||
|
||||
def test_expired_copilot_api_key_with_no_access_token_is_refused_before_any_auth_host_call(scene: _Scene) -> None:
|
||||
dl.write_copilot_key(scene.rig.copilot_dir, dl.COPILOT_STORED_KEY, scene.rig.peers.api.url, -3600)
|
||||
served: Final = _serve(scene.call(_Cell("copilot", "responses", False, "httpx")))
|
||||
assert served.status == 400, served.text
|
||||
assert dl.COPILOT_REFUSAL in _error_message(served.text), served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
|
||||
|
||||
_UNUSABLE_CHATGPT_FILES: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"invalid-json": "{not json",
|
||||
"json-list": "[]",
|
||||
"empty-file": "",
|
||||
"empty-object": "{}",
|
||||
"access-token-int": json.dumps({"access_token": 12345}),
|
||||
"access-token-list": json.dumps({"access_token": ["a", "b"]}),
|
||||
"access-token-empty": json.dumps({"access_token": ""}),
|
||||
"access-token-5kb": json.dumps({"access_token": "x" * 5120}),
|
||||
"opaque-token-without-expiry": json.dumps({"access_token": "opaque-token"}),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", tuple(_UNUSABLE_CHATGPT_FILES))
|
||||
def test_unusable_chatgpt_auth_file_is_refused_like_a_missing_one_each_time(scene: _Scene, name: str) -> None:
|
||||
(scene.rig.chatgpt_dir / "auth.json").write_text(_UNUSABLE_CHATGPT_FILES[name])
|
||||
first: Final = _serve(scene.call(_Cell("chatgpt", "responses", False, "httpx")))
|
||||
second: Final = _serve(scene.call(_Cell("chatgpt", "responses", False, "httpx")))
|
||||
for served in (first, second):
|
||||
assert served.status == 400, served.text
|
||||
assert dl.CHATGPT_REFUSAL in _error_message(served.text), served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
assert (scene.rig.chatgpt_dir / "auth.json").read_text() == _UNUSABLE_CHATGPT_FILES[name]
|
||||
|
||||
|
||||
def test_chatgpt_token_with_a_text_expiry_falls_back_to_the_expiry_inside_the_token(scene: _Scene) -> None:
|
||||
dl.write_chatgpt(scene.rig.chatgpt_dir, dl.chatgpt_record(dl.CHATGPT_STORED, expires_at="soon"))
|
||||
call: Final = scene.call(_Cell("chatgpt", "responses", False, "httpx"))
|
||||
served: Final = _serve(call)
|
||||
assert served.status == 200, served.text
|
||||
(forwarded,) = _carrying(scene.rig.peers.api.drain(), call.marker)
|
||||
assert dl.bearer(forwarded) == dl.CHATGPT_STORED
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
|
||||
|
||||
_PROVIDER_RESOLUTION_ERROR: Final = "GetLLMProvider Exception"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _UnusableCopilotFiles:
|
||||
files: Mapping[str, str]
|
||||
answer: str
|
||||
|
||||
|
||||
_UNUSABLE_COPILOT_FILES: Final[Mapping[str, _UnusableCopilotFiles]] = MappingProxyType(
|
||||
{
|
||||
"empty-access-token": _UnusableCopilotFiles(MappingProxyType({"access-token": ""}), dl.COPILOT_REFUSAL),
|
||||
"blank-access-token": _UnusableCopilotFiles(MappingProxyType({"access-token": " \n"}), dl.COPILOT_REFUSAL),
|
||||
"invalid-api-key-json": _UnusableCopilotFiles(
|
||||
MappingProxyType({"api-key.json": "{not json"}), dl.COPILOT_REFUSAL
|
||||
),
|
||||
"api-key-list": _UnusableCopilotFiles(MappingProxyType({"api-key.json": "[]"}), _PROVIDER_RESOLUTION_ERROR),
|
||||
"api-key-expiry-text": _UnusableCopilotFiles(
|
||||
MappingProxyType({"api-key.json": json.dumps({"token": "copilot-text-expiry", "expires_at": "soon"})}),
|
||||
_PROVIDER_RESOLUTION_ERROR,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", tuple(_UNUSABLE_COPILOT_FILES))
|
||||
def test_unusable_copilot_token_files_are_refused_like_missing_ones_each_time(scene: _Scene, name: str) -> None:
|
||||
variant: Final = _UNUSABLE_COPILOT_FILES[name]
|
||||
for file_name, content in variant.files.items():
|
||||
(scene.rig.copilot_dir / file_name).write_text(content)
|
||||
first: Final = _serve(scene.call(_Cell("copilot", "chat", False, "httpx")))
|
||||
second: Final = _serve(scene.call(_Cell("copilot", "chat", False, "httpx")))
|
||||
for served in (first, second):
|
||||
assert served.status == 400, served.text
|
||||
assert variant.answer in _error_message(served.text), served.text
|
||||
assert dl.COPILOT_USER_CODE not in served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", _PROVIDERS)
|
||||
def test_a_token_mounted_after_a_refusal_serves_without_a_restart_and_removing_it_refuses_again(
|
||||
scene: _Scene, provider: Provider
|
||||
) -> None:
|
||||
cell: Final = _Cell(provider, "chat", False, "httpx")
|
||||
before: Final = _serve(scene.call(cell))
|
||||
assert before.status == 400, before.text
|
||||
scene.store_valid_token(provider)
|
||||
mounted: Final = scene.call(cell)
|
||||
served: Final = _serve(mounted)
|
||||
assert served.status == 200, served.text
|
||||
assert set(rv.MARKER.findall(served.text)) == {mounted.marker}, served.text
|
||||
dl.clear_tokens(scene.rig.chatgpt_dir, scene.rig.copilot_dir)
|
||||
after: Final = _serve(scene.call(cell))
|
||||
assert after.status == 400, after.text
|
||||
assert _REFUSALS[provider] in _error_message(after.text), after.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
|
||||
|
||||
_COOLDOWN_SECONDS_LEFT: Final = 45
|
||||
|
||||
|
||||
def test_a_recent_device_code_request_does_not_hold_the_worker_for_the_rest_of_the_cooldown(scene: _Scene) -> None:
|
||||
dl.write_chatgpt(scene.rig.chatgpt_dir, {"device_code_requested_at": time.time() - (300 - _COOLDOWN_SECONDS_LEFT)})
|
||||
started: Final = time.monotonic()
|
||||
served: Final = _serve(scene.call(_Cell("chatgpt", "chat", False, "httpx")))
|
||||
elapsed: Final = time.monotonic() - started
|
||||
assert served.status == 400, served.text
|
||||
assert dl.CHATGPT_REFUSAL in _error_message(served.text), served.text
|
||||
assert elapsed < _COOLDOWN_SECONDS_LEFT - 15, elapsed
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
|
||||
|
||||
_NOT_LIVE: Final = re.compile(r"saved to the database, but the model id\(s\) \['([0-9a-f-]+)'\] are not live")
|
||||
|
||||
|
||||
def test_model_created_without_a_token_is_saved_not_live_and_serves_once_the_token_is_mounted(scene: _Scene) -> None:
|
||||
gateway: Final = scene.rig.owned.gateway
|
||||
name: Final = f"device-login-db-{uuid.uuid4().hex}"
|
||||
identity: Final = str(uuid.uuid4())
|
||||
created: Final = gateway.request(
|
||||
"POST",
|
||||
"/model/new",
|
||||
{
|
||||
"model_name": name,
|
||||
"litellm_params": {"model": "chatgpt/gpt-5.6-terra"},
|
||||
"model_info": {"id": identity},
|
||||
},
|
||||
)
|
||||
try:
|
||||
assert created.status_code == 500, created.text
|
||||
not_live: Final = _NOT_LIVE.search(created.text)
|
||||
assert not_live is not None and not_live.group(1) == identity, created.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
scene.store_valid_token("chatgpt")
|
||||
call: Final = _Call(
|
||||
scene.rig.base_url(), scene.key, name, "responses", False, "httpx", uuid.uuid4().hex, SimpleQueue()
|
||||
)
|
||||
served: Final = eventually(lambda: _serve(call), lambda answer: answer.status == 200, seconds=90)
|
||||
assert set(rv.MARKER.findall(served.text)) == {call.marker}, served.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
finally:
|
||||
gateway.request("POST", "/model/delete", {"id": identity})
|
||||
|
||||
|
||||
async def _send(client: httpx.AsyncClient, call: _Call) -> _Served:
|
||||
async with client.stream(
|
||||
"POST",
|
||||
_path(call.endpoint),
|
||||
json=_raw_body(call),
|
||||
headers={"Authorization": f"Bearer {call.key}", "anthropic-version": "2023-06-01"},
|
||||
) as response:
|
||||
raw: Final = await response.aread()
|
||||
return _record(call, _Served(response.status_code, raw.decode(), response.headers.get(_CALL_ID, "")))
|
||||
|
||||
|
||||
async def _burst(calls: tuple[_Call, ...]) -> tuple[_Served, ...]:
|
||||
async with httpx.AsyncClient(base_url=calls[0].base_url, timeout=_CLIENT_SECONDS, trust_env=False) as client:
|
||||
return tuple(await asyncio.gather(*(_send(client, call) for call in calls)))
|
||||
|
||||
|
||||
def _chat_response_id(call: _Call, served: _Served) -> str:
|
||||
if not call.stream:
|
||||
return str(rv.JSON_OBJECT.validate_json(served.text)["id"])
|
||||
first: Final = next(line for line in served.text.splitlines() if line.startswith("data: {"))
|
||||
return str(rv.JSON_OBJECT.validate_json(first[6:])["id"])
|
||||
|
||||
|
||||
_HTTPX_CELLS: Final = tuple(cell for cell in _CELLS if cell.client == "httpx")
|
||||
|
||||
|
||||
def test_a_burst_of_refusals_leaves_other_models_served_and_logs_every_request_once(scene: _Scene) -> None:
|
||||
refused_cells: Final = tuple(_HTTPX_CELLS[index % len(_HTTPX_CELLS)] for index in range(30))
|
||||
refused_calls: Final = tuple(scene.call(cell) for cell in refused_cells)
|
||||
control_calls: Final = tuple(scene.control(stream=index % 2 == 1) for index in range(10))
|
||||
answers: Final = asyncio.run(_burst((*refused_calls, *control_calls)))
|
||||
refused: Final = answers[:30]
|
||||
controls: Final = answers[30:]
|
||||
for cell, answer in zip(refused_cells, refused, strict=True):
|
||||
assert answer.status == 400, answer.text
|
||||
assert _REFUSALS[cell.provider] in _error_message(answer.text), answer.text
|
||||
for call, answer in zip(control_calls, controls, strict=True):
|
||||
assert answer.status == 200, answer.text
|
||||
assert set(rv.MARKER.findall(answer.text)) == {call.marker}, answer.text
|
||||
assert scene.rig.peers.auth_connections() == ()
|
||||
received: Final = scene.rig.peers.api.drain()
|
||||
assert sorted(rv.newest_marker(request.body.decode()) or "" for request in received) == sorted(
|
||||
call.marker for call in control_calls
|
||||
), _describe(received)
|
||||
assert {dl.bearer(request) for request in received} == {dl.CONTROL_KEY}
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE api_key = %s',
|
||||
(hashlib.sha256(scene.key.encode()).hexdigest(),),
|
||||
),
|
||||
lambda found: len(found) >= 40,
|
||||
seconds=70,
|
||||
)
|
||||
by_status: Final = {str(row["request_id"]): str(row["status"]) for row in rows}
|
||||
assert len(by_status) == len(rows) == 40, rows
|
||||
for answer in refused:
|
||||
assert by_status.get(answer.call_id) == "failure", (answer.call_id, rows)
|
||||
for call, answer in zip(control_calls, controls, strict=True):
|
||||
(logged,) = [
|
||||
request_id for request_id in by_status if rv.same_response(request_id, _chat_response_id(call, answer))
|
||||
]
|
||||
assert by_status[logged] == "success", rows
|
||||
|
||||
|
||||
def test_auth_host_outage_refuses_each_refresh_and_the_first_request_after_it_refreshes_once(scene: _Scene) -> None:
|
||||
dl.write_chatgpt(scene.rig.chatgpt_dir, dl.chatgpt_record(dl.CHATGPT_EXPIRED, refresh_token=dl.GOOD_REFRESH))
|
||||
chatgpt_cells: Final = tuple(cell for cell in _HTTPX_CELLS if cell.provider == "chatgpt")
|
||||
scene.rig.peers.switches.refuse_connect.set()
|
||||
during: Final = asyncio.run(_burst(tuple(scene.call(chatgpt_cells[index % 6]) for index in range(12))))
|
||||
for answer in during:
|
||||
assert answer.status == 400, answer.text
|
||||
assert dl.CHATGPT_REFUSAL in _error_message(answer.text), answer.text
|
||||
for answer in during:
|
||||
_failure_logged(answer)
|
||||
assert set(scene.rig.peers.auth_connections()) == {f"{dl.CHATGPT_AUTH_HOST}:443"}
|
||||
assert _describe(scene.rig.peers.auth.drain()) == ()
|
||||
assert _describe(scene.rig.peers.api.drain()) == ()
|
||||
scene.rig.peers.switches.refuse_connect.clear()
|
||||
first: Final = scene.call(chatgpt_cells[0])
|
||||
recovered: Final = _serve(first)
|
||||
assert recovered.status == 200, recovered.text
|
||||
(refresh,) = scene.rig.peers.auth.drain()
|
||||
assert (refresh.method, refresh.target) == ("POST", "/oauth/token")
|
||||
after_calls: Final = tuple(scene.call(chatgpt_cells[index % 6]) for index in range(12))
|
||||
after: Final = asyncio.run(_burst(after_calls))
|
||||
for call, answer in zip(after_calls, after, strict=True):
|
||||
assert answer.status == 200, answer.text
|
||||
assert set(rv.MARKER.findall(answer.text)) == {call.marker}, answer.text
|
||||
assert _describe(scene.rig.peers.auth.drain()) == ()
|
||||
forwarded: Final = scene.rig.peers.api.drain()
|
||||
assert {dl.bearer(request) for request in forwarded} == {dl.CHATGPT_REFRESHED}, _describe(forwarded)
|
||||
|
||||
|
||||
def _refresh_through_a_stalled_auth_host(scene: _Scene, stall: threading.Event) -> tuple[dl.Hangup, _Served, _Call]:
|
||||
dl.write_chatgpt(scene.rig.chatgpt_dir, dl.chatgpt_record(dl.CHATGPT_EXPIRED, refresh_token=dl.GOOD_REFRESH))
|
||||
call: Final = scene.call(_Cell("chatgpt", "responses", False, "httpx"))
|
||||
stall.set()
|
||||
with ThreadPoolExecutor(max_workers=1) as pool:
|
||||
pending: Final = pool.submit(_httpx, call, 240)
|
||||
try:
|
||||
dropped: Final = eventually(scene.rig.peers.dropped, lambda hangups: len(hangups) >= 1, seconds=70)
|
||||
finally:
|
||||
scene.rig.peers.switches.refuse_connect.set()
|
||||
stall.clear()
|
||||
served: Final = pending.result(timeout=240)
|
||||
assert served.status == 400, served.text
|
||||
_failure_logged(served)
|
||||
scene.rig.peers.switches.refuse_connect.clear()
|
||||
assert set(scene.rig.peers.auth_connections()) <= {f"{dl.CHATGPT_AUTH_HOST}:443"}
|
||||
return dropped[0], served, call
|
||||
|
||||
|
||||
def _assert_refused_then_refreshed(scene: _Scene, served: _Served) -> None:
|
||||
assert served.status == 400, served.text
|
||||
assert dl.CHATGPT_REFUSAL in _error_message(served.text), served.text
|
||||
assert scene.rig.peers.dropped() == ()
|
||||
assert _serve(scene.control(stream=False)).status == 200
|
||||
recovered: Final = scene.call(_Cell("chatgpt", "responses", False, "httpx"))
|
||||
answer: Final = _serve(recovered)
|
||||
assert answer.status == 200, answer.text
|
||||
assert set(rv.MARKER.findall(answer.text)) == {recovered.marker}, answer.text
|
||||
refreshes: Final = [(request.method, request.target) for request in scene.rig.peers.auth.drain()]
|
||||
assert refreshes == [("POST", "/oauth/token")], refreshes
|
||||
(forwarded,) = _carrying(scene.rig.peers.api.drain(), recovered.marker)
|
||||
assert dl.bearer(forwarded) == dl.CHATGPT_REFRESHED
|
||||
|
||||
|
||||
def test_refresh_call_gives_up_on_a_silent_auth_host_after_30_seconds_and_the_worker_recovers(scene: _Scene) -> None:
|
||||
dropped, served, _ = _refresh_through_a_stalled_auth_host(scene, scene.rig.peers.switches.hold_connect)
|
||||
assert dropped.authority == f"{dl.CHATGPT_AUTH_HOST}:443"
|
||||
assert 25 <= dropped.seconds <= 45, dropped
|
||||
_assert_refused_then_refreshed(scene, served)
|
||||
|
||||
|
||||
def test_refresh_call_gives_up_on_a_stalled_tls_handshake_at_the_connect_timeout(scene: _Scene) -> None:
|
||||
dropped, served, _ = _refresh_through_a_stalled_auth_host(scene, scene.rig.peers.switches.stall_tls)
|
||||
assert dropped.authority == f"{dl.CHATGPT_AUTH_HOST}:443"
|
||||
assert 3 <= dropped.seconds <= 20, dropped
|
||||
_assert_refused_then_refreshed(scene, served)
|
||||
230
tests/integration/sdk/test_device_code_login_guard_sdk.py
Normal file
230
tests/integration/sdk/test_device_code_login_guard_sdk.py
Normal file
|
|
@ -0,0 +1,230 @@
|
|||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import pytest
|
||||
from integration._support import device_login as dl
|
||||
from integration._support import responses_vendor as rv
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(300)
|
||||
|
||||
Provider: TypeAlias = Literal["chatgpt", "copilot"]
|
||||
|
||||
_MODELS: Final[Mapping[Provider, str]] = MappingProxyType(
|
||||
{"chatgpt": "chatgpt/gpt-5.6-terra", "copilot": "github_copilot/gpt-5.2"}
|
||||
)
|
||||
_REFUSALS: Final[Mapping[Provider, str]] = MappingProxyType(
|
||||
{"chatgpt": dl.CHATGPT_REFUSAL, "copilot": dl.COPILOT_REFUSAL}
|
||||
)
|
||||
_USER_CODES: Final[Mapping[Provider, str]] = MappingProxyType(
|
||||
{"chatgpt": dl.CHATGPT_USER_CODE, "copilot": dl.COPILOT_USER_CODE}
|
||||
)
|
||||
_AUTHENTICATORS: Final[Mapping[Provider, str]] = MappingProxyType(
|
||||
{"chatgpt": "litellm.llms.chatgpt.authenticator", "copilot": "litellm.llms.github_copilot.authenticator"}
|
||||
)
|
||||
_SCRIPT_SECONDS: Final = 60
|
||||
|
||||
_SCRIPT: Final = textwrap.dedent(
|
||||
"""
|
||||
import asyncio, json, sys, threading
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
||||
variant, model, marker, authenticator_module = sys.argv[1:5]
|
||||
messages = [{"role": "user", "content": f"Reply to marker-{marker}"}]
|
||||
|
||||
def sync_completion():
|
||||
return litellm.completion(model=model, messages=messages, num_retries=0)
|
||||
|
||||
def run():
|
||||
if variant == "completion":
|
||||
return sync_completion()
|
||||
if variant == "responses":
|
||||
return litellm.responses(model=model, input=f"Reply to marker-{marker}", num_retries=0)
|
||||
if variant == "acompletion":
|
||||
return asyncio.run(litellm.acompletion(model=model, messages=messages, num_retries=0))
|
||||
if variant == "aresponses":
|
||||
return asyncio.run(litellm.aresponses(model=model, input=f"Reply to marker-{marker}", num_retries=0))
|
||||
if variant == "worker-thread":
|
||||
outcome = {}
|
||||
def target():
|
||||
try:
|
||||
outcome["value"] = sync_completion()
|
||||
except Exception as error:
|
||||
outcome["error"] = error
|
||||
thread = threading.Thread(target=target)
|
||||
thread.start()
|
||||
thread.join()
|
||||
if "error" in outcome:
|
||||
raise outcome["error"]
|
||||
return outcome["value"]
|
||||
if variant == "batch-completion":
|
||||
(only,) = litellm.batch_completion(model=model, messages=[messages], num_retries=0)
|
||||
if isinstance(only, Exception):
|
||||
raise only
|
||||
return only
|
||||
if variant == "sync-inside-loop":
|
||||
async def main():
|
||||
return sync_completion()
|
||||
return asyncio.run(main())
|
||||
if variant == "one-liner-inside-loop":
|
||||
module = __import__(authenticator_module, fromlist=["Authenticator"])
|
||||
async def main():
|
||||
return module.Authenticator().get_access_token()
|
||||
return asyncio.run(main())
|
||||
if variant == "router-built-in-async":
|
||||
async def main():
|
||||
router = Router(model_list=[{"model_name": "login", "litellm_params": {"model": model}}])
|
||||
return await router.acompletion(model="login", messages=messages, num_retries=0)
|
||||
return asyncio.run(main())
|
||||
if variant == "wildcard-router-in-loop":
|
||||
router = Router(model_list=[{"model_name": "*", "litellm_params": {"model": "*"}}])
|
||||
async def main():
|
||||
return await router.acompletion(model=model, messages=messages, num_retries=0)
|
||||
return asyncio.run(main())
|
||||
raise SystemExit(f"unknown variant {variant}")
|
||||
|
||||
def render(result):
|
||||
if hasattr(result, "__next__"):
|
||||
return "".join(str(chunk) for chunk in result)
|
||||
return str(result)
|
||||
|
||||
try:
|
||||
print(json.dumps({"ok": True, "text": render(run())}), flush=True)
|
||||
except Exception as error:
|
||||
print(json.dumps({"ok": False, "error": f"{type(error).__name__}: {error}"}), flush=True)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Outcome:
|
||||
returncode: int
|
||||
stdout: str
|
||||
stderr: str
|
||||
|
||||
def verdict(self) -> Mapping[str, JsonValue]:
|
||||
lines: Final = [line for line in self.stdout.splitlines() if line.startswith("{")]
|
||||
assert lines, (self.stdout, self.stderr)
|
||||
return rv.JSON_OBJECT.validate_json(lines[-1])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Bench:
|
||||
peers: dl.Peers
|
||||
chatgpt_dir: Path
|
||||
copilot_dir: Path
|
||||
|
||||
def run(self, variant: str, provider: Provider, marker: str) -> _Outcome:
|
||||
completed: Final = subprocess.run(
|
||||
[sys.executable, "-P", "-c", _SCRIPT, variant, _MODELS[provider], marker, _AUTHENTICATORS[provider]],
|
||||
env={
|
||||
**os.environ,
|
||||
**self.peers.environment(),
|
||||
"CHATGPT_TOKEN_DIR": str(self.chatgpt_dir),
|
||||
"GITHUB_COPILOT_TOKEN_DIR": str(self.copilot_dir),
|
||||
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
|
||||
},
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=_SCRIPT_SECONDS,
|
||||
check=False,
|
||||
)
|
||||
return _Outcome(completed.returncode, completed.stdout, completed.stderr)
|
||||
|
||||
def token_files(self) -> tuple[str, ...]:
|
||||
return tuple(sorted(path.name for path in (*self.chatgpt_dir.iterdir(), *self.copilot_dir.iterdir())))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def bench(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Bench]:
|
||||
directory: Final = tmp_path_factory.mktemp("device-login-sdk")
|
||||
chatgpt_dir: Final = directory / "chatgpt"
|
||||
copilot_dir: Final = directory / "copilot"
|
||||
chatgpt_dir.mkdir()
|
||||
copilot_dir.mkdir()
|
||||
with dl.device_login_peers(directory) as peers:
|
||||
yield _Bench(peers, chatgpt_dir, copilot_dir)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def clean(bench: _Bench) -> Iterator[_Bench]:
|
||||
bench.peers.reset()
|
||||
dl.clear_tokens(bench.chatgpt_dir, bench.copilot_dir)
|
||||
yield bench
|
||||
bench.peers.reset()
|
||||
dl.clear_tokens(bench.chatgpt_dir, bench.copilot_dir)
|
||||
|
||||
|
||||
_FIRST_LOGIN_BEARERS: Final[Mapping[Provider, str]] = MappingProxyType(
|
||||
{"chatgpt": dl.CHATGPT_FIRST_LOGIN, "copilot": dl.COPILOT_MINTED_KEY}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("variant", ("completion", "responses"))
|
||||
@pytest.mark.parametrize("provider", ("chatgpt", "copilot"))
|
||||
def test_sync_call_on_the_main_thread_keeps_the_interactive_first_login(
|
||||
clean: _Bench, provider: Provider, variant: str
|
||||
) -> None:
|
||||
clean.peers.switches.grant.set()
|
||||
marker: Final = uuid.uuid4().hex
|
||||
outcome: Final = clean.run(variant, provider, marker)
|
||||
verdict: Final = outcome.verdict()
|
||||
assert verdict["ok"] is True, (verdict, outcome.stderr[-3000:])
|
||||
assert marker in str(verdict["text"]), verdict
|
||||
assert _USER_CODES[provider] in outcome.stdout, outcome.stdout
|
||||
assert f"{dl.CHATGPT_AUTH_HOST if provider == 'chatgpt' else dl.GITHUB_HOST}:443" in clean.peers.auth_connections()
|
||||
forwarded: Final = [request for request in clean.peers.api.drain() if request.method == "POST"]
|
||||
assert {dl.bearer(request) for request in forwarded} == {_FIRST_LOGIN_BEARERS[provider]}, [
|
||||
(request.method, request.target) for request in forwarded
|
||||
]
|
||||
expected_files: Final = ("auth.json",) if provider == "chatgpt" else ("access-token", "api-key.json")
|
||||
assert clean.token_files() == expected_files
|
||||
|
||||
|
||||
_IN_LOOP_VARIANTS: Final = (
|
||||
"acompletion",
|
||||
"aresponses",
|
||||
"worker-thread",
|
||||
"batch-completion",
|
||||
"sync-inside-loop",
|
||||
"one-liner-inside-loop",
|
||||
"router-built-in-async",
|
||||
"wildcard-router-in-loop",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("variant", _IN_LOOP_VARIANTS)
|
||||
@pytest.mark.parametrize("provider", ("chatgpt", "copilot"))
|
||||
def test_a_call_inside_a_loop_or_worker_thread_is_refused_before_any_auth_host_call(
|
||||
clean: _Bench, provider: Provider, variant: str
|
||||
) -> None:
|
||||
clean.peers.switches.grant.set()
|
||||
outcome: Final = clean.run(variant, provider, uuid.uuid4().hex)
|
||||
verdict: Final = outcome.verdict()
|
||||
assert verdict["ok"] is False, (verdict, outcome.stderr[-3000:])
|
||||
assert _REFUSALS[provider] in str(verdict["error"]), verdict
|
||||
assert _USER_CODES[provider] not in outcome.stdout, outcome.stdout
|
||||
assert clean.peers.auth_connections() == ()
|
||||
assert clean.peers.api.drain() == ()
|
||||
assert clean.token_files() == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ("chatgpt", "copilot"))
|
||||
def test_sync_call_on_the_main_thread_fails_without_hanging_when_the_auth_host_denies(
|
||||
clean: _Bench, provider: Provider
|
||||
) -> None:
|
||||
outcome: Final = clean.run("completion", provider, uuid.uuid4().hex)
|
||||
verdict: Final = outcome.verdict()
|
||||
assert verdict["ok"] is False, (verdict, outcome.stderr[-3000:])
|
||||
assert f"{dl.CHATGPT_AUTH_HOST if provider == 'chatgpt' else dl.GITHUB_HOST}:443" in clean.peers.auth_connections()
|
||||
assert clean.peers.api.drain() == ()
|
||||
assert clean.token_files() == ()
|
||||
|
|
@ -1,11 +1,18 @@
|
|||
import base64
|
||||
import json
|
||||
import time
|
||||
from unittest.mock import mock_open, patch
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.chatgpt.authenticator import Authenticator
|
||||
import litellm
|
||||
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS, HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS
|
||||
from litellm.llms.chatgpt.authenticator import (
|
||||
TOKEN_REFRESH_TIMEOUT_SECONDS,
|
||||
Authenticator,
|
||||
)
|
||||
from litellm.llms.chatgpt.common_utils import GetAccessTokenError
|
||||
|
||||
|
||||
def _make_jwt(payload: dict) -> str:
|
||||
|
|
@ -20,9 +27,9 @@ def _make_jwt(payload: dict) -> str:
|
|||
|
||||
class TestChatGPTAuthenticator:
|
||||
@pytest.fixture
|
||||
def authenticator(self):
|
||||
with patch("os.path.exists", return_value=True):
|
||||
return Authenticator()
|
||||
def authenticator(self, tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path))
|
||||
return Authenticator()
|
||||
|
||||
def test_get_access_token_from_file(self, authenticator):
|
||||
future_time = time.time() + 3600
|
||||
|
|
@ -54,10 +61,96 @@ class TestChatGPTAuthenticator:
|
|||
token = authenticator.get_access_token()
|
||||
assert token == "token-new"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("request_timeout", "expected_read"),
|
||||
[
|
||||
(DEFAULT_REQUEST_TIMEOUT_SECONDS, TOKEN_REFRESH_TIMEOUT_SECONDS),
|
||||
(TOKEN_REFRESH_TIMEOUT_SECONDS * 4.0, TOKEN_REFRESH_TIMEOUT_SECONDS),
|
||||
(TOKEN_REFRESH_TIMEOUT_SECONDS / 3.0, TOKEN_REFRESH_TIMEOUT_SECONDS / 3.0),
|
||||
],
|
||||
)
|
||||
def test_refresh_tokens_uses_bounded_timeout(self, authenticator, monkeypatch, request_timeout, expected_read):
|
||||
monkeypatch.setattr(litellm, "request_timeout", request_timeout)
|
||||
monkeypatch.setattr(litellm, "request_timeout_explicitly_set", False)
|
||||
client = MagicMock()
|
||||
response = MagicMock()
|
||||
response.json.return_value = {
|
||||
"access_token": "token-new",
|
||||
"id_token": "id-123",
|
||||
}
|
||||
client.post.return_value = response
|
||||
|
||||
with patch( # test-quality-ok: requested seam for asserting timeout propagation
|
||||
"litellm.llms.chatgpt.authenticator._get_httpx_client", return_value=client
|
||||
):
|
||||
refreshed = authenticator._refresh_tokens("refresh-123")
|
||||
|
||||
assert refreshed["access_token"] == "token-new"
|
||||
timeout = client.post.call_args.kwargs["timeout"]
|
||||
assert timeout.read == expected_read
|
||||
assert timeout.connect == HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_access_token_refuses_device_code_login_in_event_loop(self, authenticator):
|
||||
with (
|
||||
patch("builtins.open", side_effect=FileNotFoundError),
|
||||
patch.object(authenticator, "_login_device_code") as mock_login,
|
||||
patch.object(authenticator, "_wait_for_access_token") as mock_wait,
|
||||
):
|
||||
with pytest.raises(GetAccessTokenError) as exc:
|
||||
authenticator.get_access_token()
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert "event loop" in str(exc.value)
|
||||
assert authenticator.auth_file not in str(exc.value)
|
||||
mock_login.assert_not_called()
|
||||
mock_wait.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_access_token_refuses_cooldown_wait_in_event_loop(self, authenticator):
|
||||
auth_data = json.dumps({"device_code_requested_at": time.time()})
|
||||
|
||||
with (
|
||||
patch("builtins.open", mock_open(read_data=auth_data)),
|
||||
patch.object(authenticator, "_login_device_code") as mock_login,
|
||||
patch.object(authenticator, "_wait_for_access_token") as mock_wait,
|
||||
):
|
||||
with pytest.raises(GetAccessTokenError) as exc:
|
||||
authenticator.get_access_token()
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert "event loop" in str(exc.value)
|
||||
assert authenticator.auth_file not in str(exc.value)
|
||||
mock_login.assert_not_called()
|
||||
mock_wait.assert_not_called()
|
||||
|
||||
def test_get_access_token_refuses_device_code_login_in_worker_thread(self, authenticator):
|
||||
with (
|
||||
patch("builtins.open", side_effect=FileNotFoundError),
|
||||
patch.object(authenticator, "_login_device_code") as mock_login,
|
||||
patch.object(authenticator, "_wait_for_access_token") as mock_wait,
|
||||
):
|
||||
with ThreadPoolExecutor(max_workers=1) as pool:
|
||||
with pytest.raises(GetAccessTokenError) as exc:
|
||||
pool.submit(authenticator.get_access_token).result()
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert "worker thread" in str(exc.value)
|
||||
assert authenticator.auth_file not in str(exc.value)
|
||||
mock_login.assert_not_called()
|
||||
mock_wait.assert_not_called()
|
||||
|
||||
def test_get_access_token_device_code_login_without_event_loop(self, authenticator):
|
||||
with (
|
||||
patch("builtins.open", side_effect=FileNotFoundError),
|
||||
patch.object(authenticator, "_login_device_code", return_value={"access_token": "tok"}),
|
||||
):
|
||||
token = authenticator.get_access_token()
|
||||
|
||||
assert token == "tok"
|
||||
|
||||
def test_get_account_id_from_id_token(self, authenticator):
|
||||
id_token = _make_jwt(
|
||||
{"https://api.openai.com/auth": {"chatgpt_account_id": "acct-123"}}
|
||||
)
|
||||
id_token = _make_jwt({"https://api.openai.com/auth": {"chatgpt_account_id": "acct-123"}})
|
||||
auth_data = json.dumps({"id_token": id_token})
|
||||
|
||||
with (
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import json
|
||||
import os
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime, timedelta
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
|
|
@ -89,6 +90,34 @@ class TestGitHubCopilotAuthenticator:
|
|||
assert token == mock_token
|
||||
authenticator._login.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_access_token_refuses_device_code_login_in_event_loop(self, authenticator):
|
||||
with (
|
||||
patch("builtins.open", side_effect=FileNotFoundError),
|
||||
patch.object(authenticator, "_login") as mock_login,
|
||||
):
|
||||
with pytest.raises(GetAccessTokenError) as exc:
|
||||
authenticator.get_access_token()
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert "event loop" in str(exc.value)
|
||||
assert authenticator.access_token_file not in str(exc.value)
|
||||
mock_login.assert_not_called()
|
||||
|
||||
def test_get_access_token_refuses_device_code_login_in_worker_thread(self, authenticator):
|
||||
with (
|
||||
patch("builtins.open", side_effect=FileNotFoundError),
|
||||
patch.object(authenticator, "_login") as mock_login,
|
||||
):
|
||||
with ThreadPoolExecutor(max_workers=1) as pool:
|
||||
with pytest.raises(GetAccessTokenError) as exc:
|
||||
pool.submit(authenticator.get_access_token).result()
|
||||
|
||||
assert exc.value.status_code == 401
|
||||
assert "worker thread" in str(exc.value)
|
||||
assert authenticator.access_token_file not in str(exc.value)
|
||||
mock_login.assert_not_called()
|
||||
|
||||
def test_get_access_token_failure(self, authenticator):
|
||||
"""Test that an exception is raised after multiple login failures."""
|
||||
with (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue