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:
devin-ai-integration[bot] 2026-10-06 03:20:07 +00:00 • committed by GitHub
parent 64cddd6e13
commit ab3a59fe81
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 2128 additions and 8 deletions

View file

@ -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.

View file

@ -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())

View file

@ -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:

View 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()

View 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

View 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)

View 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() == ()

View file

@ -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 (

View file

@ -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 (