diff --git a/lm15/_authlock.py b/lm15/_authlock.py index ffa6d757..4e316e0b 100644 --- a/lm15/_authlock.py +++ b/lm15/_authlock.py @@ -36,7 +36,7 @@ from pathlib import Path from typing import Any, Iterator -from .errors import LockTimeoutError +from .errors import LockTimeoutError, UnsupportedFeatureError _DEFAULT_LOCK_TIMEOUT_S = 60.0 _LOCK_POLL_INTERVAL_S = 0.05 @@ -73,10 +73,18 @@ def lock_path_for(path: Path) -> Path: return _lock_dir() / f"{digest}.lock" +# Chosen by which primitive the platform actually has, not by `os.name`: WASI reports +# "posix" and ships no fcntl, so a name test picks an implementation that cannot import. try: import fcntl -except ImportError: # pragma: no cover - Pyodide (and any POSIX build without it) - fcntl = None # type: ignore[assignment] +except ImportError: # pragma: no cover - Windows, and POSIX-ish builds without fcntl + fcntl = None + +try: + import msvcrt +except ImportError: # pragma: no cover - every non-Windows platform + msvcrt = None + if fcntl is not None: @@ -90,19 +98,7 @@ def _try_lock(fd: int) -> bool: def _unlock(fd: int) -> None: fcntl.flock(fd, fcntl.LOCK_UN) -elif os.name == "posix": # pragma: no cover - Pyodide: the lock is asked for, not merely imported - - def _try_lock(fd: int) -> bool: - raise CredentialLockTimeout( - "shared credential locking needs fcntl, which this Python (Pyodide?) does not have; " - "pass an explicit credential instead of a stored login" - ) - - def _unlock(fd: int) -> None: - return None - -else: # pragma: no cover - exercised only on Windows - import msvcrt +elif msvcrt is not None: # pragma: no cover - exercised only on Windows def _try_lock(fd: int) -> bool: try: @@ -117,6 +113,18 @@ def _unlock(fd: int) -> None: except OSError: pass +else: # pragma: no cover - platforms with neither primitive, such as WASI + + def _try_lock(fd: int) -> bool: + raise UnsupportedFeatureError( + "This platform provides no advisory file locking (neither fcntl nor msvcrt), " + "so lm15 cannot serialize credential refreshes against other processes. " + "Reading credentials still works; refreshing them from here does not." + ) + + def _unlock(fd: int) -> None: + return None + @contextmanager def hold_file_lock( diff --git a/lm15/auth.py b/lm15/auth.py index 981e6305..fe6565dc 100644 --- a/lm15/auth.py +++ b/lm15/auth.py @@ -84,12 +84,27 @@ "write_xai_credential", ] -CLAUDE_CODE_CREDENTIALS_PATH = Path("~/.claude/.credentials.json").expanduser() +def _user_path(*parts: str) -> Path: + """Resolve a path under the user's home without raising where there is none. + + Evaluated at import to define the credential-path constants below, so it must not + raise. `Path.expanduser()` raises RuntimeError wherever no home directory can be + determined -- a WASI guest, some container setups -- and these constants are + defaults a caller can always override, so an unexpanded path is a better answer + than a failed import. Anything that actually opens it still fails, and says why. + """ + try: + return Path("~", *parts).expanduser() + except RuntimeError: + return Path("~", *parts) + + +CLAUDE_CODE_CREDENTIALS_PATH = _user_path(".claude", ".credentials.json") CLAUDE_CODE_CLIENT_ID = "9d1c250a-e61b-44d5-88ed-5944d1962f5e" CLAUDE_CODE_TOKEN_URL = "https://platform.claude.com/v1/oauth/token" CLAUDE_CODE_LOGIN_HINT = "Log in again: run `claude` and use /login (Claude subscription auth)" -CODEX_CLI_AUTH_PATH = Path("~/.codex/auth.json").expanduser() +CODEX_CLI_AUTH_PATH = _user_path(".codex", "auth.json") OPENAI_CODEX_CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann" OPENAI_CODEX_TOKEN_URL = "https://auth.openai.com/oauth/token" OPENAI_CODEX_JWT_CLAIM_PATH = "https://api.openai.com/auth" @@ -506,7 +521,7 @@ def get_codex_cli_access_token( XAI_TOKEN_URL = "https://auth.x.ai/oauth2/token" XAI_OAUTH_SCOPE = "openid profile email offline_access grok-cli:access api:access" XAI_LOGIN_HINT = "Log in again: run lm15.auth.login_xai() (SuperGrok / X Premium subscription auth)" -PI_AGENT_AUTH_PATH = Path("~/.pi/agent/auth.json").expanduser() +PI_AGENT_AUTH_PATH = _user_path(".pi", "agent", "auth.json") _XAI_PROVIDER_KEY = "xai" _XAI_DEFAULT_TOKEN_LIFETIME_S = 3600 diff --git a/lm15/transports/_async.py b/lm15/transports/_async.py index baea91df..478aa5db 100644 --- a/lm15/transports/_async.py +++ b/lm15/transports/_async.py @@ -11,13 +11,16 @@ CancelledError without awaiting anything that could itself be cancelled. - Timeouts use a cancellation-safe wait helper, including on Python 3.10/3.11 where asyncio.wait_for can swallow simultaneous caller cancellation. +- TLS lives in `_ssl.py`, which this module imports on the first https request + and not before. That module needs the stdlib `ssl`; this one does not, so a + CPython build without `ssl` still carries plain HTTP through here. """ from __future__ import annotations import asyncio import select import socket -from typing import AsyncIterator +from typing import TYPE_CHECKING, AsyncIterator from ._exceptions import ( ConnectError, @@ -34,11 +37,13 @@ build_request_head, ) from ._proxy import ProxyRoute, connect_payload, proxy_route_for, route_origin -from ._ssl import SSLError, make_ssl_context from ._timeouts import wait_for from ._types import AsyncTransportResponse, TransportRequest from ._url import ParsedURL, parse_url +if TYPE_CHECKING: + from ._ssl import TLS + _DEFAULT_CONNECT_TIMEOUT = 10.0 _DEFAULT_READ_TIMEOUT = 60.0 @@ -196,16 +201,18 @@ def __init__( self._ca_bundle = ca_bundle self._proxy = proxy self._trust_env = trust_env - self._ssl_ctx: ssl.SSLContext | None = None + self._tls: TLS | None = None self._pool = _AsyncConnectionPool(max_connections) self._closed = False - def _get_ssl_ctx(self) -> ssl.SSLContext: - if self._ssl_ctx is None: - self._ssl_ctx = make_ssl_context( - verify=self._verify, ca_bundle=self._ca_bundle - ) - return self._ssl_ctx + def _tls_half(self) -> "TLS": + """Build the TLS half on first https use. `._ssl` needs the stdlib `ssl` module, + which a reduced build may not have — and a plain-HTTP caller never asks for it.""" + from ._ssl import TLS + + if self._tls is None: + self._tls = TLS(verify=self._verify, ca_bundle=self._ca_bundle) + return self._tls def pool_stats(self) -> dict: return self._pool.stats() @@ -355,55 +362,42 @@ async def _open( origin: tuple[str, str, int], connect_timeout: float, ) -> _AsyncConnection: - ctx = self._get_ssl_ctx() if parsed.is_tls else None + tls = self._tls_half() if parsed.is_tls else None - if proxy is not None and parsed.is_tls: + if proxy is not None and tls is not None: # CONNECT over a raw socket, then hand it to open_connection # for the end-to-end TLS handshake (works on 3.10; StreamWriter # gained start_tls only in 3.11). tunnel = await self._connect_tunnel(parsed, proxy, timeout=connect_timeout) + reader, writer = await tls.connect_over( + tunnel, server_hostname=parsed.host, timeout=connect_timeout + ) + return _AsyncConnection(origin, reader, writer) + + connect_host = proxy.host if proxy is not None else parsed.host + connect_port = proxy.port if proxy is not None else parsed.port + if tls is not None: + reader, writer = await tls.connect( + host=connect_host, + port=connect_port, + server_hostname=parsed.host, + timeout=connect_timeout, + ) + else: try: reader, writer = await wait_for( - asyncio.open_connection( - sock=tunnel, ssl=ctx, server_hostname=parsed.host - ), + asyncio.open_connection(host=connect_host, port=connect_port), timeout=connect_timeout, cancel_result=lambda pair: pair[1].close(), ) except asyncio.TimeoutError as exc: - tunnel.close() - raise ConnectTimeout("TLS handshake through proxy timed out") from exc - except SSLError as exc: - tunnel.close() - raise ConnectError(f"TLS handshake failed: {exc}") from exc + raise ConnectTimeout( + f"timed out connecting to {connect_host}:{connect_port}" + ) from exc except OSError as exc: - tunnel.close() - raise ConnectError(f"TLS handshake through proxy failed: {exc}") from exc - return _AsyncConnection(origin, reader, writer) - - connect_host = proxy.host if proxy is not None else parsed.host - connect_port = proxy.port if proxy is not None else parsed.port - try: - reader, writer = await wait_for( - asyncio.open_connection( - host=connect_host, - port=connect_port, - ssl=ctx, - server_hostname=parsed.host if ctx else None, - ), - timeout=connect_timeout, - cancel_result=lambda pair: pair[1].close(), - ) - except asyncio.TimeoutError as exc: - raise ConnectTimeout( - f"timed out connecting to {connect_host}:{connect_port}" - ) from exc - except SSLError as exc: - raise ConnectError(f"TLS handshake failed: {exc}") from exc - except OSError as exc: - raise ConnectError( - f"failed to connect to {connect_host}:{connect_port}: {exc}" - ) from exc + raise ConnectError( + f"failed to connect to {connect_host}:{connect_port}: {exc}" + ) from exc # TCP_NODELAY on the underlying socket sock = writer.get_extra_info("socket") diff --git a/lm15/transports/_ssl.py b/lm15/transports/_ssl.py index d0c2d7ee..b9ed1062 100644 --- a/lm15/transports/_ssl.py +++ b/lm15/transports/_ssl.py @@ -1,44 +1,110 @@ -"""SSL context factory. +"""TLS for the stdlib transports: one context, and the ways to put it on a socket. -We rely on the stdlib `ssl` module's `create_default_context`, which on -Python 3.10+ loads the system trust store correctly on Linux/macOS/Windows. -No certifi bundle is shipped — set SSL_CERT_FILE if your system store is -broken, or pass an explicit ca_bundle= to the transport. +The socket transports load this module only when a request names an HTTPS URL. +This keeps plain HTTP usable on Python builds without ``ssl``. If such a build +does request HTTPS, :class:`TLS` raises a transport error that names the host's +fetch transport as the Pyodide alternative. -Under Pyodide the `ssl` module is absent (there is no socket to wrap); -importing lm15 must still work there — the fetch transport carries the -wire — so the import is optional and the socket transports refuse by -name at connect time, not at import time. +We rely on the stdlib ``create_default_context``, which on Python 3.10+ loads +the system trust store correctly on Linux, macOS, and Windows. No certifi bundle +is shipped: set ``SSL_CERT_FILE`` if the system store is broken, or pass an +explicit ``ca_bundle=`` to the transport. """ from __future__ import annotations +import asyncio +import socket + try: import ssl -except ImportError: # pragma: no cover - Pyodide +except ImportError: # pragma: no cover - Pyodide and reduced Python builds ssl = None # type: ignore[assignment] -if ssl is not None: - SSLError = ssl.SSLError -else: # pragma: no cover - Pyodide - - class SSLError(OSError): - """Never raised where there is no ssl; keeps the socket transports' except clauses valid.""" - - -def make_ssl_context( - *, verify: bool = True, ca_bundle: str | None = None -) -> "ssl.SSLContext": - if ssl is None: # pragma: no cover - Pyodide - from ._exceptions import ConnectError - - raise ConnectError( - "this Python has no ssl module (Pyodide?), so the socket transports cannot open TLS; " - "use lm15.transports.FetchTransport, the host's fetch" - ) - if not verify: - ctx = ssl._create_unverified_context() - return ctx - ctx = ssl.create_default_context() - if ca_bundle: - ctx.load_verify_locations(cafile=ca_bundle) - return ctx +from ._exceptions import ConnectError, ConnectTimeout +from ._timeouts import wait_for + + +def _close(sock: socket.socket) -> None: + try: + sock.close() + except Exception: + pass + + +def _is_ssl_error(exc: OSError) -> bool: + return ssl is not None and isinstance(exc, ssl.SSLError) + + +class TLS: + """The TLS half of a transport, holding the context its connections share.""" + + __slots__ = ("_ctx",) + + def __init__(self, *, verify: bool = True, ca_bundle: str | None = None) -> None: + if ssl is None: # pragma: no cover - Pyodide + raise ConnectError( + "this Python has no ssl module (Pyodide?), so the socket transports " + "cannot open TLS; use lm15.transports.FetchTransport, the host's fetch" + ) + if not verify: + self._ctx = ssl._create_unverified_context() + return + self._ctx = ssl.create_default_context() + if ca_bundle: + self._ctx.load_verify_locations(cafile=ca_bundle) + + def wrap( + self, sock: socket.socket, *, server_hostname: str, timeout: float + ) -> socket.socket: + """Return ``sock`` with TLS, or close it and describe the handshake failure.""" + try: + sock.settimeout(timeout) + return self._ctx.wrap_socket(sock, server_hostname=server_hostname) + except OSError as exc: + _close(sock) + raise ConnectError(f"TLS handshake failed: {exc}") from exc + + async def connect( + self, *, host: str, port: int, server_hostname: str, timeout: float + ) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]: + """Open a TCP connection and run its TLS handshake in one asyncio call.""" + try: + return await wait_for( + asyncio.open_connection( + host=host, + port=port, + ssl=self._ctx, + server_hostname=server_hostname, + ), + timeout=timeout, + cancel_result=lambda pair: pair[1].close(), + ) + except asyncio.TimeoutError as exc: + raise ConnectTimeout(f"timed out connecting to {host}:{port}") from exc + except OSError as exc: + if _is_ssl_error(exc): + raise ConnectError(f"TLS handshake failed: {exc}") from exc + raise ConnectError(f"failed to connect to {host}:{port}: {exc}") from exc + + async def connect_over( + self, sock: socket.socket, *, server_hostname: str, timeout: float + ) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]: + """Run TLS over a socket that a proxy has already tunneled to the origin.""" + try: + return await wait_for( + asyncio.open_connection( + sock=sock, + ssl=self._ctx, + server_hostname=server_hostname, + ), + timeout=timeout, + cancel_result=lambda pair: pair[1].close(), + ) + except asyncio.TimeoutError as exc: + _close(sock) + raise ConnectTimeout("TLS handshake through proxy timed out") from exc + except OSError as exc: + _close(sock) + if _is_ssl_error(exc): + raise ConnectError(f"TLS handshake failed: {exc}") from exc + raise ConnectError(f"TLS handshake through proxy failed: {exc}") from exc diff --git a/lm15/transports/_sync.py b/lm15/transports/_sync.py index e0a2bc6d..c7c5434c 100644 --- a/lm15/transports/_sync.py +++ b/lm15/transports/_sync.py @@ -1,5 +1,5 @@ """ -Sync transport built on the stdlib `socket` + `ssl` modules. +Sync transport built on the stdlib `socket` module. Design: @@ -15,6 +15,9 @@ socket, readable means the server sent EOF — so we drop it and open fresh. - If the server sent `Connection: close`, or the body is still outstanding when the response closes, the connection is closed rather than reused. +- TLS lives in `_ssl.py`, which this module imports on the first https request + and not before. That module needs the stdlib `ssl`; this one does not, so a + CPython build without `ssl` still carries plain HTTP through here. - Proxies: `proxy=` pins one explicitly; otherwise `trust_env=True` (default) honors HTTP_PROXY / HTTPS_PROXY / ALL_PROXY / NO_PROXY. Plain-HTTP targets are forwarded absolute-URI; TLS targets are tunneled with CONNECT and the @@ -25,7 +28,7 @@ import select import socket import threading -from typing import Iterator +from typing import TYPE_CHECKING, Iterator from ._exceptions import ( ConnectError, @@ -45,10 +48,12 @@ build_request_head, ) from ._proxy import ProxyRoute, connect_payload, proxy_route_for, route_origin -from ._ssl import SSLError, make_ssl_context from ._types import TransportRequest, TransportResponse from ._url import ParsedURL, parse_url +if TYPE_CHECKING: + from ._ssl import TLS + _DEFAULT_CONNECT_TIMEOUT = 10.0 _DEFAULT_READ_TIMEOUT = 60.0 @@ -207,18 +212,20 @@ def __init__( self._ca_bundle = ca_bundle self._proxy = proxy self._trust_env = trust_env - self._ssl_ctx: ssl.SSLContext | None = None - self._ssl_lock = threading.Lock() + self._tls: TLS | None = None + self._tls_lock = threading.Lock() self._pool = _ConnectionPool(max_connections) self._closed = False - def _get_ssl_ctx(self) -> ssl.SSLContext: - with self._ssl_lock: - if self._ssl_ctx is None: - self._ssl_ctx = make_ssl_context( - verify=self._verify, ca_bundle=self._ca_bundle - ) - return self._ssl_ctx + def _tls_half(self) -> "TLS": + """Build the TLS half on first https use. `._ssl` needs the stdlib `ssl` module, + which a reduced build may not have — and a plain-HTTP caller never asks for it.""" + from ._ssl import TLS + + with self._tls_lock: + if self._tls is None: + self._tls = TLS(verify=self._verify, ca_bundle=self._ca_bundle) + return self._tls def pool_stats(self) -> dict: return self._pool.stats() @@ -336,6 +343,10 @@ def _open( origin: tuple[str, str, int], connect_timeout: float, ) -> _SyncConnection: + # Built before the socket, as on the async side: a missing `ssl` or an unreadable + # ca_bundle then fails with nothing yet open. + tls = self._tls_half() if parsed.is_tls else None + connect_host = proxy.host if proxy is not None else parsed.host connect_port = proxy.port if proxy is not None else parsed.port try: @@ -366,17 +377,8 @@ def _open( pass raise - if parsed.is_tls: - try: - ctx = self._get_ssl_ctx() - sock.settimeout(connect_timeout) - sock = ctx.wrap_socket(sock, server_hostname=parsed.host) - except (SSLError, OSError) as exc: - try: - sock.close() - except Exception: - pass - raise ConnectError(f"TLS handshake failed: {exc}") from exc + if tls is not None: + sock = tls.wrap(sock, server_hostname=parsed.host, timeout=connect_timeout) return _SyncConnection(origin, sock)