diff --git a/litellm/litellm_core_utils/asyncify.py b/litellm/litellm_core_utils/asyncify.py index b58e707b8f8..17cfb3ef84f 100644 --- a/litellm/litellm_core_utils/asyncify.py +++ b/litellm/litellm_core_utils/asyncify.py @@ -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. diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 563826c2b93..5ab16cd189a 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -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()) diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index 80fd4f755e7..7821756bc16 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -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: diff --git a/tests/integration/_support/device_login.py b/tests/integration/_support/device_login.py new file mode 100644 index 00000000000..9b990fd2167 --- /dev/null +++ b/tests/integration/_support/device_login.py @@ -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() diff --git a/tests/integration/providers/test_device_code_login_guard_boot.py b/tests/integration/providers/test_device_code_login_guard_boot.py new file mode 100644 index 00000000000..0c5b2615d6f --- /dev/null +++ b/tests/integration/providers/test_device_code_login_guard_boot.py @@ -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 diff --git a/tests/integration/providers/test_device_code_login_guard_wire.py b/tests/integration/providers/test_device_code_login_guard_wire.py new file mode 100644 index 00000000000..8ec41ac6ac2 --- /dev/null +++ b/tests/integration/providers/test_device_code_login_guard_wire.py @@ -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) diff --git a/tests/integration/sdk/test_device_code_login_guard_sdk.py b/tests/integration/sdk/test_device_code_login_guard_sdk.py new file mode 100644 index 00000000000..d579fbef237 --- /dev/null +++ b/tests/integration/sdk/test_device_code_login_guard_sdk.py @@ -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() == () diff --git a/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py b/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py index a9ced2afcf9..d367e4df0ab 100644 --- a/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py +++ b/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py @@ -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 ( diff --git a/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py b/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py index 6c846a90c71..a49a4b44b74 100644 --- a/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py +++ b/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py @@ -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 (