diff --git a/solstone/convey/secure_listener/accept.py b/solstone/convey/secure_listener/accept.py index d64ac8d08..466d9e61a 100644 --- a/solstone/convey/secure_listener/accept.py +++ b/solstone/convey/secure_listener/accept.py @@ -12,15 +12,15 @@ import logging import socket import uuid from collections.abc import Callable -from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field -from typing import Any, Literal +from typing import Any, Final, Literal from OpenSSL import SSL from solstone.think.link.auth import AuthorizedClients from solstone.think.link.window import window_open +from .admission import SecureListenerAdmission from .framing import RESET_INTERNAL_ERROR from .identity import ConveyIdentity from .mux import RESET_CTX_NO_IDENTITY, Multiplexer, ResetDiagnostic, StreamWriter @@ -34,8 +34,13 @@ CERTLESS_TUNNEL_CAP = 4 CERTLESS_PAIR_FAILURE_CAP = 3 CERTLESS_INVALID_NONCE_BACKOFF_SECONDS = 1.0 CERTLESS_WINDOW_POLL_SECONDS = 5.0 +SECURE_LISTENER_TCP_DRAIN_TIMEOUT_SECONDS: Final[float] = 120.0 +KEEPALIVE_IDLE_SECONDS: Final[int] = 30 +KEEPALIVE_INTERVAL_SECONDS: Final[int] = 10 +KEEPALIVE_COUNT: Final[int] = 3 PeerMode = Literal["pl-direct", "pl-via-spl"] +QueuedFrame = tuple[int, int, bytes] def certless_admission_mode( @@ -84,7 +89,7 @@ class SecureListener: strict_tls_ctx: SSL.Context, relaxed_tls_ctx: SSL.Context, authorized: AuthorizedClients, - executor: ThreadPoolExecutor, + admission: SecureListenerAdmission, callosum_emit: CallosumEmit | None = None, host: str = "0.0.0.0", port: int = 7657, @@ -94,7 +99,7 @@ class SecureListener: self._strict_tls_ctx = strict_tls_ctx self._relaxed_tls_ctx = relaxed_tls_ctx self._authorized = authorized - self._executor = executor + self._admission = admission self._emit = callosum_emit or (lambda _event, _fields: None) self._host = host self._port = port @@ -180,6 +185,7 @@ class SecureListener: ) -> None: connection_id = uuid.uuid4().hex mode = _mode_from_peername(writer.get_extra_info("peername")) + _apply_tcp_keepalive(writer, self._log) self._log.info( "secure connection accepted conn=%s mode=%s", connection_id, @@ -234,7 +240,8 @@ class SecureListener: admitted_certless = certless_mode is not None tls_ctx = self._relaxed_tls_ctx if admitted_certless else self._strict_tls_ctx tls = new_server(tls_ctx) - send_queue: asyncio.Queue[bytes] = asyncio.Queue() + send_queue: asyncio.PriorityQueue[QueuedFrame] = asyncio.PriorityQueue() + send_sequence = 0 identity: ConveyIdentity | None = None certless_handle: CertlessConnection | None = None certless_registered = False @@ -243,10 +250,16 @@ class SecureListener: if not data: return tcp_writer.write(data) - await tcp_writer.drain() + await asyncio.wait_for( + tcp_writer.drain(), + timeout=SECURE_LISTENER_TCP_DRAIN_TIMEOUT_SECONDS, + ) - async def send_frame(frame: bytes) -> None: - send_queue.put_nowait(frame) + async def send_frame(frame: bytes, *, urgent: bool = False) -> None: + nonlocal send_sequence + priority = 0 if urgent else 1 + send_queue.put_nowait((priority, send_sequence, frame)) + send_sequence += 1 async def handle_stream( reader: asyncio.StreamReader, @@ -272,7 +285,7 @@ class SecureListener: reader, writer, asyncio.get_running_loop(), - self._executor, + self._admission, ) close_after_failure = await self._record_certless_dispatch( certless_handle, @@ -293,7 +306,7 @@ class SecureListener: reader, writer, asyncio.get_running_loop(), - self._executor, + self._admission, ) finally: self._log.debug( @@ -389,7 +402,7 @@ class SecureListener: async def app_writer_loop() -> None: while True: - data = await send_queue.get() + _priority, _sequence, data = await send_queue.get() try: outbound = _encrypt(tls, data) await write_ciphertext(outbound) @@ -405,13 +418,19 @@ class SecureListener: name=f"secure-tcp-writer-{connection_id}", ) try: - await reader_task + done, pending = await asyncio.wait( + {reader_task, writer_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + for task in done: + task.result() finally: if certless_registered: self._unregister_certless_connection(connection_id) - writer_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await writer_task + for task in (reader_task, writer_task): + if not task.done(): + task.cancel() + await asyncio.gather(reader_task, writer_task, return_exceptions=True) await mux.close() def _identity_for_peer(self, mode: PeerMode, fingerprint: str) -> ConveyIdentity: @@ -499,11 +518,11 @@ def _encrypt(tls: Any, plaintext: bytes) -> bytes: async def _drain_send_queue( tls: Any, write_ciphertext: Callable[[bytes], asyncio.Future[Any] | Any], - queue: asyncio.Queue[bytes], + queue: asyncio.PriorityQueue[QueuedFrame], ) -> None: while True: try: - chunk = queue.get_nowait() + _priority, _sequence, chunk = queue.get_nowait() except asyncio.QueueEmpty: break try: @@ -536,4 +555,47 @@ def _tls_close_reason(exc: TlsError) -> str: return "tls_alert" +def _tcp_keepalive_options(socket_module: Any = socket) -> list[tuple[int, int, int]]: + options: list[tuple[int, int, int]] = [] + if hasattr(socket_module, "SOL_SOCKET") and hasattr(socket_module, "SO_KEEPALIVE"): + options.append((socket_module.SOL_SOCKET, socket_module.SO_KEEPALIVE, 1)) + if hasattr(socket_module, "IPPROTO_TCP"): + tcp_level = socket_module.IPPROTO_TCP + if hasattr(socket_module, "TCP_KEEPIDLE"): + options.append( + (tcp_level, socket_module.TCP_KEEPIDLE, KEEPALIVE_IDLE_SECONDS) + ) + if hasattr(socket_module, "TCP_KEEPALIVE"): + options.append( + (tcp_level, socket_module.TCP_KEEPALIVE, KEEPALIVE_IDLE_SECONDS) + ) + if hasattr(socket_module, "TCP_KEEPINTVL"): + options.append( + (tcp_level, socket_module.TCP_KEEPINTVL, KEEPALIVE_INTERVAL_SECONDS) + ) + if hasattr(socket_module, "TCP_KEEPCNT"): + options.append((tcp_level, socket_module.TCP_KEEPCNT, KEEPALIVE_COUNT)) + return options + + +def _apply_tcp_keepalive( + tcp_writer: asyncio.StreamWriter, + logger: logging.Logger, +) -> None: + sock = tcp_writer.get_extra_info("socket") + if sock is None or not hasattr(sock, "setsockopt"): + return + for level, optname, value in _tcp_keepalive_options(): + try: + sock.setsockopt(level, optname, value) + except Exception: + logger.warning( + "secure listener TCP keepalive option failed level=%r opt=%r value=%r", + level, + optname, + value, + exc_info=True, + ) + + __all__ = ["SecureListener", "certless_admission_mode"] diff --git a/solstone/convey/secure_listener/admission.py b/solstone/convey/secure_listener/admission.py new file mode 100644 index 000000000..feea83860 --- /dev/null +++ b/solstone/convey/secure_listener/admission.py @@ -0,0 +1,411 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +"""In-process admission control for secure-listener WSGI work.""" + +from __future__ import annotations + +import asyncio +import logging +import threading +import time +from collections import deque +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from typing import Any, Final, Literal + +from solstone.think.utils import get_config + +log = logging.getLogger("convey.secure_listener.admission") + +DEFAULT_SECURE_LISTENER_CAPACITY: Final[int] = 16 +DEFAULT_SECURE_LISTENER_STREAMING_CAPACITY: Final[int] = 8 +MAX_SECURE_LISTENER_CAPACITY: Final[int] = 128 + +_PermitKind = Literal["total", "streaming", "streaming_over_budget"] + + +class SecureListenerAdmissionRejected(Exception): + """The listener refused admission before executor submission.""" + + +@dataclass(frozen=True) +class SecureListenerAdmissionConfig: + capacity: int = DEFAULT_SECURE_LISTENER_CAPACITY + streaming_capacity: int = DEFAULT_SECURE_LISTENER_STREAMING_CAPACITY + refuse_when_full: bool = False + + @property + def queue_limit(self) -> int: + return self.capacity * 2 + + +class SecureListenerPermit: + """One in-process serving-capacity permit.""" + + def __init__( + self, + admission: SecureListenerAdmission, + kind: _PermitKind, + *, + queue_wait_ms: float = 0.0, + ) -> None: + self._admission = admission + self.kind = kind + self.queue_wait_ms = queue_wait_ms + self._released = False + self._lock = threading.Lock() + + def release(self) -> None: + with self._lock: + if self._released: + return + self._released = True + self._admission._release(self.kind) + + def __enter__(self) -> SecureListenerPermit: + return self + + def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> None: + self.release() + + +@dataclass +class _QueuedWaiter: + loop: asyncio.AbstractEventLoop + future: asyncio.Future[SecureListenerPermit] + queued_at: float + permit: SecureListenerPermit | None = None + + +class _PermitHandoff: + def __init__(self, permit: SecureListenerPermit) -> None: + self._permit = permit + self._started = False + self._released = False + self._lock = threading.Lock() + + def mark_started(self) -> bool: + with self._lock: + if self._released: + return False + self._started = True + return True + + def release_from_worker(self) -> None: + permit = self._release(started_required=True) + if permit is not None: + permit.release() + + def release_from_caller_if_not_started(self) -> bool: + permit = self._release(started_required=False) + if permit is None: + return False + permit.release() + return True + + def _release(self, *, started_required: bool) -> SecureListenerPermit | None: + with self._lock: + if self._released: + return None + if started_required and not self._started: + return None + if not started_required and self._started: + return None + self._released = True + return self._permit + + +def resolve_admission_config() -> SecureListenerAdmissionConfig: + cfg = get_config() + link_cfg = cfg.get("link") if isinstance(cfg, dict) else None + if not isinstance(link_cfg, dict): + return SecureListenerAdmissionConfig() + + capacity = _resolve_int( + link_cfg, + "secure_listener_capacity", + default=DEFAULT_SECURE_LISTENER_CAPACITY, + valid_min=1, + valid_max=MAX_SECURE_LISTENER_CAPACITY, + warning_default="16", + ) + streaming_default = min(DEFAULT_SECURE_LISTENER_STREAMING_CAPACITY, capacity) + streaming_capacity = _resolve_int( + link_cfg, + "secure_listener_streaming_capacity", + default=streaming_default, + valid_min=0, + valid_max=capacity, + warning_default=str(streaming_default), + ) + refuse_when_full = _resolve_bool( + link_cfg, + "secure_listener_refuse_when_full", + default=False, + warning_default="false", + ) + return SecureListenerAdmissionConfig( + capacity=capacity, + streaming_capacity=streaming_capacity, + refuse_when_full=refuse_when_full, + ) + + +def _resolve_int( + link_cfg: dict[str, Any], + key: str, + *, + default: int, + valid_min: int, + valid_max: int, + warning_default: str, +) -> int: + raw = link_cfg.get(key, default) + if ( + not isinstance(raw, int) + or isinstance(raw, bool) + or raw < valid_min + or raw > valid_max + ): + log.warning( + "Invalid link.%s in journal config: %r \u2014 defaulting to %s", + key, + raw, + warning_default, + ) + return default + return raw + + +def _resolve_bool( + link_cfg: dict[str, Any], + key: str, + *, + default: bool, + warning_default: str, +) -> bool: + raw = link_cfg.get(key, default) + if not isinstance(raw, bool): + log.warning( + "Invalid link.%s in journal config: %r \u2014 defaulting to %s", + key, + raw, + warning_default, + ) + return default + return raw + + +class SecureListenerAdmission: + """Admission, queueing, and content-free telemetry for listener work.""" + + def __init__( + self, + config: SecureListenerAdmissionConfig | None = None, + *, + thread_name_prefix: str = "secure-listener-wsgi", + ) -> None: + self.config = config or SecureListenerAdmissionConfig() + self._executor = ThreadPoolExecutor( + max_workers=self.config.capacity, + thread_name_prefix=thread_name_prefix, + ) + self._lock = threading.Lock() + self._streaming_condition = threading.Condition(self._lock) + self._waiters: deque[_QueuedWaiter] = deque() + self._active_total = 0 + self._active_streaming = 0 + self._active_streaming_over_budget = 0 + self._rejected_total = 0 + self._rejected_streaming = 0 + self._admitted_streaming_over_budget = 0 + + async def acquire(self) -> SecureListenerPermit: + started = time.monotonic() + loop = asyncio.get_running_loop() + with self._lock: + if self._active_total < self.config.capacity and not self._waiters: + self._active_total += 1 + return SecureListenerPermit(self, "total") + if ( + self.config.refuse_when_full + and len(self._waiters) >= self.config.queue_limit + ): + self._rejected_total += 1 + raise SecureListenerAdmissionRejected + waiter = _QueuedWaiter( + loop=loop, + future=loop.create_future(), + queued_at=started, + ) + self._waiters.append(waiter) + + try: + return await waiter.future + except BaseException: + if self._cancel_waiter(waiter): + raise + if waiter.permit is not None: + waiter.permit.release() + raise + + async def submit( + self, + loop: asyncio.AbstractEventLoop, + func: Callable[..., Any], + *args: Any, + ) -> Any: + permit = await self.acquire() + handoff = _PermitHandoff(permit) + try: + future = loop.run_in_executor( + self._executor, + self._run_with_permit, + handoff, + func, + args, + ) + except BaseException: + handoff.release_from_caller_if_not_started() + raise + try: + return await future + except asyncio.CancelledError: + future.cancel() + handoff.release_from_caller_if_not_started() + raise + + def try_acquire_streaming(self) -> SecureListenerPermit | None: + if self.config.streaming_capacity <= 0: + return None + with self._lock: + if self._active_streaming >= self.config.streaming_capacity: + return None + self._active_streaming += 1 + return SecureListenerPermit(self, "streaming") + + def acquire_streaming( + self, + timeout_s: float, + *, + cancel_event: threading.Event, + ) -> SecureListenerPermit | None: + if self.config.streaming_capacity <= 0: + return None + deadline = time.monotonic() + max(0.0, timeout_s) + with self._streaming_condition: + while self._active_streaming >= self.config.streaming_capacity: + if cancel_event.is_set(): + return None + remaining = deadline - time.monotonic() + if remaining <= 0: + return None + self._streaming_condition.wait(remaining) + self._active_streaming += 1 + return SecureListenerPermit(self, "streaming") + + def reject_streaming(self) -> None: + with self._lock: + self._rejected_streaming += 1 + + def admit_streaming_over_budget(self) -> SecureListenerPermit: + with self._lock: + self._active_streaming_over_budget += 1 + self._admitted_streaming_over_budget += 1 + return SecureListenerPermit(self, "streaming_over_budget") + + def snapshot(self) -> dict[str, Any]: + now = time.monotonic() + with self._lock: + longest_wait_ms = 0.0 + if self._waiters: + longest_wait_ms = max( + (now - waiter.queued_at) * 1000.0 for waiter in self._waiters + ) + return { + "timestamp_ms": int(time.time() * 1000), + "active": { + "total": self._active_total, + "streaming": self._active_streaming, + "streaming_over_budget": self._active_streaming_over_budget, + }, + "limit": { + "total": self.config.capacity, + "streaming": self.config.streaming_capacity, + "queue": self.config.queue_limit, + }, + "queued": {"total": len(self._waiters)}, + "rejected": { + "total": self._rejected_total, + "streaming": self._rejected_streaming, + }, + "admitted_over_budget": { + "streaming": self._admitted_streaming_over_budget, + }, + "longest_wait_ms": longest_wait_ms, + "refusal_enabled": self.config.refuse_when_full, + } + + def shutdown(self, *, wait: bool = True, cancel_futures: bool = True) -> None: + self._executor.shutdown(wait=wait, cancel_futures=cancel_futures) + + def _run_with_permit( + self, + handoff: _PermitHandoff, + func: Callable[..., Any], + args: tuple[Any, ...], + ) -> Any: + if not handoff.mark_started(): + raise SecureListenerAdmissionRejected + try: + return func(*args) + finally: + handoff.release_from_worker() + + def _release(self, kind: _PermitKind) -> None: + with self._streaming_condition: + if kind == "total": + self._active_total = max(0, self._active_total - 1) + self._wake_waiters_locked() + elif kind == "streaming": + self._active_streaming = max(0, self._active_streaming - 1) + self._streaming_condition.notify() + elif kind == "streaming_over_budget": + self._active_streaming_over_budget = max( + 0, + self._active_streaming_over_budget - 1, + ) + + def _cancel_waiter(self, waiter: _QueuedWaiter) -> bool: + with self._lock: + try: + self._waiters.remove(waiter) + except ValueError: + return False + return True + + def _wake_waiters_locked(self) -> None: + while self._waiters and self._active_total < self.config.capacity: + waiter = self._waiters.popleft() + if waiter.future.cancelled(): + continue + self._active_total += 1 + permit = SecureListenerPermit( + self, + "total", + queue_wait_ms=(time.monotonic() - waiter.queued_at) * 1000.0, + ) + waiter.permit = permit + waiter.loop.call_soon_threadsafe(self._deliver_waiter, waiter, permit) + + def _deliver_waiter( + self, + waiter: _QueuedWaiter, + permit: SecureListenerPermit, + ) -> None: + if waiter.future.cancelled(): + permit.release() + return + waiter.future.set_result(permit) diff --git a/solstone/convey/secure_listener/mux.py b/solstone/convey/secure_listener/mux.py index e875e03d4..8d711cd9b 100644 --- a/solstone/convey/secure_listener/mux.py +++ b/solstone/convey/secure_listener/mux.py @@ -23,7 +23,7 @@ import asyncio import logging from collections.abc import Awaitable, Callable from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Final +from typing import Final, Protocol from .framing import ( FLAG_CLOSE, @@ -58,10 +58,12 @@ from .framing import ( validate_flags, ) -if TYPE_CHECKING: - StreamHandler = Callable[[asyncio.StreamReader, "StreamWriter"], Awaitable[None]] -else: - StreamHandler = object +StreamHandler = Callable[[asyncio.StreamReader, "StreamWriter"], Awaitable[None]] + + +class FrameSender(Protocol): + def __call__(self, frame: bytes, *, urgent: bool = False) -> Awaitable[None]: ... + RESET_CTX_MALFORMED_FRAME: Final[str] = "malformed_frame" RESET_CTX_PARITY_VIOLATION: Final[str] = "parity_violation" @@ -197,7 +199,7 @@ class Multiplexer: def __init__( self, - send_frame: Callable[[bytes], Awaitable[None]], + send_frame: FrameSender, handler: StreamHandler, *, is_listener: bool = True, @@ -395,7 +397,7 @@ class Multiplexer: await self._tunnel_fatal(RESET_CTX_MALFORMED_FRAME) return if is_ping: - await self._emit(build_pong(nonce)) + await self._emit(build_pong(nonce), urgent=True) # Stray PONG: the listener does not initiate pings (the dialer drives # keepalive), so an unsolicited PONG is silently dropped per # proto/framing.md § responder behavior. @@ -468,14 +470,14 @@ class Multiplexer: return False return (stream_id % 2 == 1) if self._is_listener else (stream_id % 2 == 0) - async def _emit(self, frame: Frame) -> None: + async def _emit(self, frame: Frame, *, urgent: bool = False) -> None: if self._closed: return try: encoded = frame.encode() except ProtocolError: return - await self._send_frame(encoded) + await self._send_frame(encoded, urgent=urgent) def _fire_diag(self, stream_id: int, reason: int, context: str) -> None: if self._on_reset is None: diff --git a/solstone/convey/secure_listener/runtime.py b/solstone/convey/secure_listener/runtime.py index da147f2e3..3b07165b6 100644 --- a/solstone/convey/secure_listener/runtime.py +++ b/solstone/convey/secure_listener/runtime.py @@ -10,7 +10,6 @@ import atexit import logging import socket import threading -from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from typing import Any @@ -21,6 +20,7 @@ from solstone.think.link.paths import LinkState, authorized_clients_path, ca_dir from solstone.think.utils import get_journal, journal_is_active from .accept import SecureListener +from .admission import SecureListenerAdmission, resolve_admission_config from .tls import build_relaxed_server_context, build_server_context, issue_server_cert logger = logging.getLogger("convey.secure_listener.runtime") @@ -33,11 +33,8 @@ class RuntimeState: started_event: threading.Event = field(default_factory=threading.Event) apps: list[Any] = field(default_factory=list) authorized: AuthorizedClients | None = None - executor: ThreadPoolExecutor = field( - default_factory=lambda: ThreadPoolExecutor( - max_workers=16, - thread_name_prefix="secure-listener-wsgi", - ) + admission: SecureListenerAdmission = field( + default_factory=lambda: SecureListenerAdmission(resolve_admission_config()) ) listener: SecureListener | None = None start_error: BaseException | None = None @@ -96,7 +93,7 @@ def _thread_main(runtime: RuntimeState) -> None: strict_tls_ctx=strict_tls_ctx, relaxed_tls_ctx=relaxed_tls_ctx, authorized=authorized, - executor=runtime.executor, + admission=runtime.admission, callosum_emit=emit, ) runtime.listener = listener @@ -120,7 +117,13 @@ def _thread_main(runtime: RuntimeState) -> None: asyncio.gather(*pending, return_exceptions=True) ) finally: - runtime.executor.shutdown(wait=True, cancel_futures=True) + loop.run_until_complete( + asyncio.to_thread( + runtime.admission.shutdown, + wait=True, + cancel_futures=True, + ) + ) loop.close() diff --git a/solstone/convey/secure_listener/wsgi.py b/solstone/convey/secure_listener/wsgi.py index 2b37935c5..28793e4e4 100644 --- a/solstone/convey/secure_listener/wsgi.py +++ b/solstone/convey/secure_listener/wsgi.py @@ -6,11 +6,12 @@ from __future__ import annotations import asyncio +import json import sys import threading import urllib.parse from collections.abc import Iterable -from concurrent.futures import CancelledError, ThreadPoolExecutor +from concurrent.futures import CancelledError from dataclasses import dataclass from typing import Any, Callable, Final @@ -18,6 +19,7 @@ from werkzeug.exceptions import HTTPException from solstone.think.link.window import window_open +from .admission import SecureListenerAdmission, SecureListenerAdmissionRejected from .identity import ConveyIdentity from .mux import ( RESET_CTX_APP_CANCELLATION, @@ -31,6 +33,9 @@ _NO_BODY_STATUSES = {204, 304} _DEFAULT_PORTS: Final[dict[str, int]] = {"http": 80, "https": 443} WSGI_SEND_BRIDGE_POLL_SECONDS: Final[float] = 0.5 WSGI_INPUT_READ_TIMEOUT_SECONDS: Final[float] = 120.0 +STREAMING_PERMIT_WAIT_TIMEOUT_SECONDS: Final[float] = 1.0 +CAPACITY_SNAPSHOT_PATH: Final[str] = "/__solstone/secure-listener/capacity" +_SAFE_STREAM_REFUSAL_METHODS: Final[frozenset[str]] = frozenset({"GET", "HEAD"}) # The cert-less pairing tunnel admits EXACTLY these endpoints (canonical + # the /app/link legacy alias). Never relax to a suffix/substring match — that # would admit pair_start and widen the tunnel. @@ -43,6 +48,12 @@ class HttpBadRequest(ValueError): ... class _WsgiClientDisconnected(ConnectionError): ... +class _WsgiResponseRefused(Exception): + def __init__(self, status: int) -> None: + super().__init__(status) + self.status = status + + def _normalize_location_headers( headers: list[tuple[str, str]], host_header: str | None, @@ -283,7 +294,7 @@ async def dispatch_stream( stream_reader: asyncio.StreamReader, stream_writer: StreamWriter, loop: asyncio.AbstractEventLoop, - executor: ThreadPoolExecutor, + admission: SecureListenerAdmission, ) -> DispatchResult: try: request = await parse_http_head(stream_reader) @@ -298,6 +309,17 @@ async def dispatch_stream( return DispatchResult(endpoint=None, status=400) stream_writer.report_recv_consumed(request.head_bytes) + if request.path == CAPACITY_SNAPSHOT_PATH: + if identity.fingerprint is not None and request.method in {"GET", "HEAD"}: + await write_json_response( + stream_writer, + 200, + "OK", + admission.snapshot(), + include_body=request.method != "HEAD", + ) + return DispatchResult(endpoint=None, status=200) + transfer_encoding = (request.transfer_encoding or "").lower() if transfer_encoding: await _finish_early( @@ -369,20 +391,29 @@ async def dispatch_stream( report_recv_consumed, path_info, ) - future = loop.run_in_executor( - executor, - _run_wsgi, - app, - environ, - stream_writer, - loop, - disconnect_event, - ) try: - status = await future + status = await admission.submit( + loop, + _run_wsgi, + app, + environ, + stream_writer, + loop, + disconnect_event, + admission, + ) wsgi_input = environ["wsgi.input"] if wsgi_input.remaining > 0: stream_writer.begin_drain(RESET_CTX_APP_CANCELLATION) + except SecureListenerAdmissionRejected: + await write_json_response( + stream_writer, + 503, + "Service Unavailable", + {"error": "secure listener capacity is full"}, + ) + stream_writer.begin_drain(RESET_CTX_BODY_DISCARD_CANCELLATION) + return DispatchResult(endpoint=endpoint, status=503) except asyncio.CancelledError: disconnect_event.set() raise @@ -475,6 +506,27 @@ async def write_simple_response( await writer.close() +async def write_json_response( + writer: StreamWriter, + status_code: int, + reason: str, + payload: dict[str, Any], + *, + include_body: bool = True, +) -> None: + body = json.dumps(payload, separators=(",", ":")).encode("utf-8") + response_body = body if include_body else b"" + head = ( + f"HTTP/1.1 {status_code} {reason}\r\n" + "Content-Type: application/json\r\n" + f"Content-Length: {len(body)}\r\n" + "Connection: close\r\n" + "\r\n" + ).encode("ascii") + await writer.write(head + response_body) + await writer.close() + + async def _finish_early( stream_writer: StreamWriter, status_code: int, @@ -493,6 +545,7 @@ def _run_wsgi( stream_writer: StreamWriter, loop: asyncio.AbstractEventLoop, disconnect_event: threading.Event, + admission: SecureListenerAdmission, ) -> int: state: dict[str, Any] = { "status": None, @@ -501,6 +554,7 @@ def _run_wsgi( "chunked": False, "body_allowed": True, } + streaming_permit = None def send(data: bytes) -> None: if disconnect_event.is_set(): @@ -545,6 +599,7 @@ def _run_wsgi( return write def ensure_head() -> None: + nonlocal streaming_permit if state["headers_sent"]: return status = state["status"] or "500 Internal Server Error" @@ -567,6 +622,34 @@ def _run_wsgi( if body_allowed and not has_content_length and not has_transfer_encoding: headers.append(("Transfer-Encoding", "chunked")) state["chunked"] = True + has_transfer_encoding = True + if _should_use_streaming_lane( + headers, + body_allowed=body_allowed, + has_content_length=has_content_length, + has_transfer_encoding=has_transfer_encoding, + admission=admission, + ): + permit = admission.try_acquire_streaming() + if permit is not None: + streaming_permit = permit + else: + method_is_safe = method in _SAFE_STREAM_REFUSAL_METHODS + if method_is_safe: + admission.reject_streaming() + _send_refusal_response(send, method) + state["headers_sent"] = True + raise _WsgiResponseRefused(503) + # why: a response that got this far on a non-safe method may + # have committed side effects, so it must queue rather than refuse. + permit = admission.acquire_streaming( + STREAMING_PERMIT_WAIT_TIMEOUT_SECONDS, + cancel_event=disconnect_event, + ) + if permit is None: + streaming_permit = admission.admit_streaming_over_budget() + else: + streaming_permit = permit lines = [f"HTTP/1.1 {status}\r\n"] lines.extend(f"{name}: {value}\r\n" for name, value in headers) lines.append("\r\n") @@ -602,9 +685,15 @@ def _run_wsgi( send(b"0\r\n\r\n") future = asyncio.run_coroutine_threadsafe(stream_writer.close(), loop) future.result() + except _WsgiResponseRefused as exc: + future = asyncio.run_coroutine_threadsafe(stream_writer.close(), loop) + future.result() + return exc.status except _WsgiClientDisconnected: return 499 finally: + if streaming_permit is not None: + streaming_permit.release() disconnect_event.set() if response_iter is not None and hasattr(response_iter, "close"): response_iter.close() # type: ignore[attr-defined] @@ -616,3 +705,38 @@ def _status_code(status: str) -> int: return int(status.split(" ", 1)[0]) except (ValueError, IndexError): return 500 + + +def _should_use_streaming_lane( + headers: list[tuple[str, str]], + *, + body_allowed: bool, + has_content_length: bool, + has_transfer_encoding: bool, + admission: SecureListenerAdmission, +) -> bool: + if not body_allowed or admission.config.streaming_capacity <= 0: + return False + content_type = "" + for name, value in headers: + if name.lower() == "content-type": + content_type = value.split(";", 1)[0].strip().lower() + break + return ( + content_type == "text/event-stream" + or has_transfer_encoding + or not has_content_length + ) + + +def _send_refusal_response(send: Callable[[bytes], None], method: str) -> None: + body = b'{"error":"secure listener streaming capacity is full"}\n' + response_body = b"" if method == "HEAD" else body + head = ( + "HTTP/1.1 503 Service Unavailable\r\n" + "Content-Type: application/json\r\n" + f"Content-Length: {len(body)}\r\n" + "Connection: close\r\n" + "\r\n" + ).encode("ascii") + send(head + response_body) diff --git a/solstone/think/link/README.md b/solstone/think/link/README.md index 7ce002069..e30c99da2 100644 --- a/solstone/think/link/README.md +++ b/solstone/think/link/README.md @@ -32,6 +32,7 @@ The `spl` repo's `home/` continues as the open-source reference implementation o TLS termination, multiplexing, and inline WSGI dispatch now live in `solstone/convey/secure_listener/`, because Convey owns both listening ports: the DL web port and the PL secure-listener port 7657. +Secure-listener capacity is configured with `link.secure_listener_capacity`; `link.secure_listener_streaming_capacity = 0` disables the streaming lane split. ## naming diff --git a/tests/link/certless_helpers.py b/tests/link/certless_helpers.py index 196767f86..e8c8d5bc8 100644 --- a/tests/link/certless_helpers.py +++ b/tests/link/certless_helpers.py @@ -8,7 +8,6 @@ import contextlib import hashlib import json import urllib.parse -from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from pathlib import Path from typing import Any, Literal @@ -20,6 +19,10 @@ from cryptography.hazmat.primitives import hashes, serialization from cryptography.hazmat.primitives.asymmetric import ec from cryptography.x509.oid import NameOID +from solstone.convey.secure_listener.admission import ( + SecureListenerAdmission, + SecureListenerAdmissionConfig, +) from solstone.convey.secure_listener.identity import ConveyIdentity from solstone.convey.secure_listener.wsgi import DispatchResult, dispatch_stream from solstone.think.link.ca import ca_pin_matches @@ -528,8 +531,17 @@ async def dispatch_request( reader.feed_eof() writer = FakeStreamWriter() loop = asyncio.get_running_loop() - with ThreadPoolExecutor(max_workers=1) as executor: - result = await dispatch_stream(app, identity, reader, writer, loop, executor) + admission = SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=1, + streaming_capacity=1, + refuse_when_full=False, + ) + ) + try: + result = await dispatch_stream(app, identity, reader, writer, loop, admission) + finally: + await asyncio.to_thread(admission.shutdown, wait=True, cancel_futures=True) status, response_headers, response_body = _parse_http_response(bytes(writer.data)) return DispatchResponse( result=result, diff --git a/tests/link/pairing_harness.py b/tests/link/pairing_harness.py index 0953476e8..d7797ead1 100644 --- a/tests/link/pairing_harness.py +++ b/tests/link/pairing_harness.py @@ -9,7 +9,6 @@ import queue import threading import time from collections.abc import Awaitable, Callable, Iterator -from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from pathlib import Path from typing import Any @@ -18,6 +17,10 @@ import pytest from OpenSSL import SSL from solstone.apps.network.routes import _build_pair_link +from solstone.convey.secure_listener.admission import ( + SecureListenerAdmission, + SecureListenerAdmissionConfig, +) from solstone.convey.secure_listener.identity import ConveyIdentity from solstone.convey.secure_listener.mux import Multiplexer, StreamWriter from solstone.convey.secure_listener.tls import ( @@ -37,6 +40,7 @@ PairingStreamHandler = Callable[ [asyncio.StreamReader, StreamWriter], Awaitable[None], ] +QueuedFrame = tuple[int, int, bytes] @dataclass @@ -48,8 +52,14 @@ class PairingHarness: handle_stream: PairingStreamHandler | None = None host: str = "127.0.0.1" port: int = 0 - _executor: ThreadPoolExecutor = field( - default_factory=lambda: ThreadPoolExecutor(max_workers=2), + _admission: SecureListenerAdmission = field( + default_factory=lambda: SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=2, + streaming_capacity=2, + refuse_when_full=False, + ) + ), ) _ready: queue.Queue[BaseException | None] = field(default_factory=queue.Queue) _loop: asyncio.AbstractEventLoop | None = None @@ -109,7 +119,7 @@ class PairingHarness: loop.call_soon_threadsafe(loop.stop) if self._thread is not None: self._thread.join(timeout=5) - self._executor.shutdown(wait=True, cancel_futures=True) + self._admission.shutdown(wait=True, cancel_futures=True) def _run_loop(self) -> None: loop = asyncio.new_event_loop() @@ -157,7 +167,8 @@ class PairingHarness: tcp_writer: asyncio.StreamWriter, ) -> None: tls = new_server(self.relaxed_ctx) - send_queue: asyncio.Queue[bytes] = asyncio.Queue() + send_queue: asyncio.PriorityQueue[QueuedFrame] = asyncio.PriorityQueue() + send_sequence = 0 loop = asyncio.get_running_loop() identity = ConveyIdentity( mode="pl-via-spl", @@ -173,8 +184,11 @@ class PairingHarness: tcp_writer.write(data) await tcp_writer.drain() - async def send_frame(frame: bytes) -> None: - send_queue.put_nowait(frame) + async def send_frame(frame: bytes, *, urgent: bool = False) -> None: + nonlocal send_sequence + priority = 0 if urgent else 1 + send_queue.put_nowait((priority, send_sequence, frame)) + send_sequence += 1 async def handle_stream( reader: asyncio.StreamReader, @@ -189,7 +203,7 @@ class PairingHarness: reader, writer, loop, - self._executor, + self._admission, ) mux = Multiplexer(send_frame, handle_stream, is_listener=True) @@ -207,7 +221,7 @@ class PairingHarness: async def writer_loop() -> None: while True: - frame = await send_queue.get() + _priority, _sequence, frame = await send_queue.get() await write_ciphertext(_encrypt(tls, frame)) reader_task = asyncio.create_task(reader_loop()) @@ -279,10 +293,10 @@ def _encrypt(tls: Any, plaintext: bytes) -> bytes: async def _drain_send_queue( tls: Any, write_ciphertext: Callable[[bytes], Awaitable[None]], - send_queue: asyncio.Queue[bytes], + send_queue: asyncio.PriorityQueue[QueuedFrame], ) -> None: - drained: list[bytes] = [] + drained: list[QueuedFrame] = [] while not send_queue.empty(): drained.append(send_queue.get_nowait()) - for frame in drained: + for _priority, _sequence, frame in drained: await write_ciphertext(_encrypt(tls, frame)) diff --git a/tests/link/secure_listener_harness.py b/tests/link/secure_listener_harness.py index df6bed02f..64742cc24 100644 --- a/tests/link/secure_listener_harness.py +++ b/tests/link/secure_listener_harness.py @@ -3,7 +3,7 @@ from __future__ import annotations -from concurrent.futures import ThreadPoolExecutor +import asyncio from dataclasses import dataclass from pathlib import Path from typing import Any @@ -11,6 +11,10 @@ from typing import Any import pytest from solstone.convey.secure_listener.accept import SecureListener +from solstone.convey.secure_listener.admission import ( + SecureListenerAdmission, + SecureListenerAdmissionConfig, +) from solstone.convey.secure_listener.tls import ( build_relaxed_server_context, build_server_context, @@ -30,7 +34,7 @@ class SecureListenerHarness: ca: LoadedCa authorized: AuthorizedClients listener: SecureListener - executor: ThreadPoolExecutor + admission: SecureListenerAdmission host: str port: int @@ -62,8 +66,12 @@ class SecureListenerHarness: server_key, authorized, ) - executor = ThreadPoolExecutor( - max_workers=4, + admission = SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=4, + streaming_capacity=4, + refuse_when_full=False, + ), thread_name_prefix="secure-listener-e2e-wsgi", ) listener = SecureListener( @@ -71,7 +79,7 @@ class SecureListenerHarness: strict_tls_ctx=strict_tls_ctx, relaxed_tls_ctx=relaxed_tls_ctx, authorized=authorized, - executor=executor, + admission=admission, callosum_emit=lambda _event, _fields: None, host="127.0.0.1", port=0, @@ -85,7 +93,7 @@ class SecureListenerHarness: ca=ca, authorized=authorized, listener=listener, - executor=executor, + admission=admission, host=str(host), port=int(port), ) @@ -94,7 +102,11 @@ class SecureListenerHarness: try: await self.listener.stop() finally: - self.executor.shutdown(wait=True, cancel_futures=True) + await asyncio.to_thread( + self.admission.shutdown, + wait=True, + cancel_futures=True, + ) def seed_nonce( self, diff --git a/tests/link/test_browser_pairing.py b/tests/link/test_browser_pairing.py index 97caed056..e1a1f5e62 100644 --- a/tests/link/test_browser_pairing.py +++ b/tests/link/test_browser_pairing.py @@ -43,7 +43,7 @@ class PairWs: raise ConnectionClosed(None, None) return self.frames.pop(0) - async def send(self, data: bytes) -> None: + async def send(self, data: bytes, *, urgent: bool = False) -> None: self.sent.append(data) async def close(self) -> None: diff --git a/tests/link/test_client_session.py b/tests/link/test_client_session.py index a0456fefc..82ff993e2 100644 --- a/tests/link/test_client_session.py +++ b/tests/link/test_client_session.py @@ -46,7 +46,7 @@ class FakeTransport: self._auto_pong = auto_pong self._decoder = FrameDecoder() - async def send(self, data: bytes) -> None: + async def send(self, data: bytes, *, urgent: bool = False) -> None: self.sent.append(data) if not self._auto_pong: return @@ -294,7 +294,7 @@ def _probing_body_source(probe: list[int]) -> client.BodySource: async def test_dialer_mux_ping_emits_matching_pong() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) nonce = b"12345678" @@ -313,7 +313,7 @@ async def test_dialer_mux_pong_records_liveness_without_emit() -> None: sent: list[bytes] = [] inbound_count = 0 - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) def on_inbound() -> None: @@ -332,7 +332,7 @@ async def test_dialer_mux_malformed_control_frame_closes_without_raising() -> No sent: list[bytes] = [] dispatched: list[Frame] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -357,7 +357,7 @@ async def test_dialer_mux_malformed_control_frame_closes_without_raising() -> No async def test_dialer_window_credit_exact_cap_is_accepted() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -373,7 +373,7 @@ async def test_dialer_window_credit_exact_cap_is_accepted() -> None: async def test_dialer_window_credit_overflow_resets_and_forgets() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -394,7 +394,7 @@ async def test_dialer_window_credit_overflow_resets_and_forgets() -> None: async def test_dialer_window_credit_uses_remaining_credit_accounting() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -458,7 +458,7 @@ async def test_dialer_window_credit_uses_remaining_credit_accounting() -> None: async def test_dialer_invalid_flags_on_known_stream_reset_and_forget() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -480,7 +480,7 @@ async def test_dialer_invalid_open_flags_close_existing_stream_before_open_polic ): sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -499,7 +499,7 @@ async def test_dialer_invalid_open_flags_close_existing_stream_before_open_polic async def test_dialer_misplaced_control_on_unknown_stream_resets() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -517,7 +517,7 @@ async def test_dialer_misplaced_control_on_unknown_stream_resets() -> None: async def test_dialer_misplaced_control_on_known_stream_beats_invalid_data() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -537,7 +537,7 @@ async def test_dialer_misplaced_control_on_known_stream_beats_invalid_data() -> async def test_dialer_sibling_stream_survives_window_overflow() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -563,7 +563,7 @@ async def test_dialer_sibling_stream_survives_window_overflow() -> None: async def test_dialer_sibling_stream_survives_invalid_flags() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -589,7 +589,7 @@ async def test_dialer_sibling_stream_survives_invalid_flags() -> None: async def test_dialer_sibling_stream_survives_misplaced_control() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -615,7 +615,7 @@ async def test_dialer_sibling_stream_survives_misplaced_control() -> None: async def test_dialer_unknown_stream_close_is_ignored() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -630,7 +630,7 @@ async def test_dialer_unknown_stream_close_is_ignored() -> None: async def test_dialer_unknown_stream_reset_is_ignored() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) @@ -645,7 +645,7 @@ async def test_dialer_unknown_stream_reset_is_ignored() -> None: async def test_dialer_unknown_stream_data_and_window_get_reset() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = client._DialerMultiplexer(send) diff --git a/tests/link/test_mux.py b/tests/link/test_mux.py index ba8955325..351165855 100644 --- a/tests/link/test_mux.py +++ b/tests/link/test_mux.py @@ -5,20 +5,25 @@ from __future__ import annotations import asyncio import contextlib +import json import logging import threading import time -from concurrent.futures import ThreadPoolExecutor +from collections.abc import Iterator from dataclasses import dataclass from pathlib import Path from typing import Any import pytest -from flask import Response, request +from flask import Response, jsonify, request from solstone.convey import root as root_module from solstone.convey.secure_listener import mux as mux_module from solstone.convey.secure_listener import wsgi as wsgi_module +from solstone.convey.secure_listener.admission import ( + SecureListenerAdmission, + SecureListenerAdmissionConfig, +) from solstone.convey.secure_listener.framing import ( FLAG_CLOSE, FLAG_DATA, @@ -74,6 +79,7 @@ from solstone.think.link.auth import AuthorizedClients from solstone.think.link.client import _http_head_bytes, _parse_http_response from solstone.think.link.paths import authorized_clients_path from tests.link.certless_helpers import ( + FakeStreamWriter, certless_identity, make_convey_app, pl_identity, @@ -104,6 +110,24 @@ def _authorize_fingerprint( monkeypatch.setattr(root_module, "get_authorized_clients", lambda: authorized) +def _admission( + capacity: int = 1, + *, + streaming_capacity: int | None = None, + refuse_when_full: bool = False, +) -> SecureListenerAdmission: + resolved_streaming_capacity = ( + capacity if streaming_capacity is None else streaming_capacity + ) + return SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=capacity, + streaming_capacity=resolved_streaming_capacity, + refuse_when_full=refuse_when_full, + ) + ) + + class _ObservedEvent(asyncio.Event): def __init__(self, waiter_entered: asyncio.Event) -> None: super().__init__() @@ -158,6 +182,45 @@ async def _wait_for_stream_data(sent: list[bytes], stream_id: int) -> bytes: raise AssertionError(f"stream {stream_id} did not emit DATA") +async def _dispatch_raw_request( + app: Any, + identity: Any, + admission: SecureListenerAdmission, + method: str, + path: str, + *, + headers: dict[str, str] | None = None, + body: bytes = b"", +) -> tuple[int, dict[str, str], bytes, FakeStreamWriter]: + reader = asyncio.StreamReader() + reader.feed_data( + _http_head_bytes( + method, + path, + headers=headers, + content_length=len(body), + ) + + body + ) + reader.feed_eof() + writer = FakeStreamWriter() + result = await dispatch_stream( + app, + identity, + reader, + writer, + asyncio.get_running_loop(), + admission, + ) + status, response_headers, response_body = _parse_http_response(bytes(writer.data)) + assert result.status == status + return status, response_headers, response_body, writer + + +async def _shutdown_admission(admission: SecureListenerAdmission) -> None: + await asyncio.to_thread(admission.shutdown, wait=True, cancel_futures=True) + + def _assert_single_reset( frames: list[Frame], stream_id: int, @@ -202,7 +265,7 @@ async def _new_parked_writer(stream_id: int = 1) -> _ParkedWriter: handler_done = asyncio.Event() captured: dict[str, Any] = {} - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: return async def handler(reader: asyncio.StreamReader, writer: Any) -> None: @@ -240,7 +303,7 @@ async def _new_parked_read_task(stream_id: int = 1) -> _OpenReader: handler_done = asyncio.Event() captured: dict[str, Any] = {} - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: return async def handler(reader: asyncio.StreamReader, writer: Any) -> None: @@ -406,7 +469,7 @@ async def test_open_with_initial_payload_hits_handler() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = Multiplexer(send, handler, is_listener=True) @@ -438,7 +501,7 @@ async def test_open_payload_over_window_is_rejected_without_opening() -> None: sent: list[bytes] = [] handler_invoked = False - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: # pragma: no cover - must not run @@ -465,7 +528,7 @@ async def test_open_payload_at_window_boundary_opens() -> None: """An OPEN payload of exactly INITIAL_WINDOW is within the window (the guard is strictly greater-than) and opens the stream normally.""" - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: return async def handler(*_: object) -> None: @@ -487,7 +550,7 @@ async def test_open_payload_at_window_boundary_opens() -> None: async def test_wrong_parity_stream_id_gets_reset() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -505,7 +568,7 @@ async def test_wrong_parity_stream_id_gets_reset() -> None: async def test_unknown_stream_data_gets_reset() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -524,7 +587,7 @@ async def test_unknown_stream_bare_close_is_ignored() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -543,7 +606,7 @@ async def test_unknown_stream_bare_reset_is_ignored() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -562,7 +625,7 @@ async def test_unknown_stream_data_still_gets_reset_with_diagnostic() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -589,7 +652,7 @@ async def test_unknown_stream_window_gets_reset() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -615,7 +678,7 @@ async def test_unknown_stream_window_gets_reset() -> None: async def test_listener_window_credit_exact_cap_is_accepted() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -640,7 +703,7 @@ async def test_listener_window_credit_overflow_resets_and_terminates() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -673,7 +736,7 @@ async def test_listener_invalid_flags_on_unknown_stream_reset_without_state() -> sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -702,7 +765,7 @@ async def test_listener_invalid_open_flags_reject_before_opening() -> None: diags: list[ResetDiagnostic] = [] handler_invoked = False - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -734,7 +797,7 @@ async def test_listener_invalid_flags_on_known_stream_terminate_without_payload( sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -766,7 +829,7 @@ async def test_listener_misplaced_pong_on_known_stream_resets_and_terminates() - sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -797,7 +860,7 @@ async def test_listener_sibling_stream_survives_window_overflow() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -835,7 +898,7 @@ async def test_listener_sibling_stream_survives_invalid_flags() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -875,7 +938,7 @@ async def test_listener_sibling_stream_survives_misplaced_control() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -923,7 +986,7 @@ async def test_concurrent_streams_do_not_interfere() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = Multiplexer(send, handler, is_listener=True) @@ -953,7 +1016,7 @@ async def test_concurrent_streams_do_not_interfere() -> None: async def test_validates_open_reopen_is_protocol_error() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) gate = asyncio.Event() @@ -980,7 +1043,7 @@ async def test_validates_open_reopen_is_protocol_error() -> None: async def test_listener_local_stream_ids_do_not_recycle_after_stream_close() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1006,7 +1069,7 @@ async def test_listener_local_stream_ids_do_not_recycle_after_stream_close() -> def test_listener_local_stream_id_exhaustion_still_raises() -> None: - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: return async def handler(*_: object) -> None: @@ -1026,7 +1089,7 @@ def test_listener_local_stream_id_exhaustion_still_raises() -> None: async def test_ping_emits_matching_pong() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1047,7 +1110,7 @@ async def test_ping_emits_matching_pong() -> None: async def test_repeated_pings_each_get_matching_pong() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1067,7 +1130,7 @@ async def test_repeated_pings_each_get_matching_pong() -> None: async def test_unsolicited_pong_is_silently_dropped() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1095,7 +1158,7 @@ async def test_stream_zero_malformed_control_raises_tunnel_fatal_once( sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1122,7 +1185,7 @@ async def test_decoder_corrupt_frame_fatal_tears_down_streams_without_diag_storm sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1159,7 +1222,7 @@ async def test_second_feed_after_tunnel_fatal_is_inert() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1185,7 +1248,7 @@ async def test_ping_on_nonzero_stream_is_protocol_error() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -1226,7 +1289,7 @@ async def test_pings_interleave_with_open_streams() -> None: sent: list[bytes] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) mux = Multiplexer(send, handler, is_listener=True) @@ -1250,6 +1313,31 @@ async def test_pings_interleave_with_open_streams() -> None: await mux.close() +@pytest.mark.asyncio +async def test_control_pong_uses_urgent_send_path() -> None: + sends: list[tuple[bool, bytes]] = [] + + async def send(data: bytes, *, urgent: bool = False) -> None: + sends.append((urgent, data)) + + async def handler(*_: object) -> None: + await asyncio.Event().wait() + + mux = Multiplexer(send, handler, is_listener=True) + nonce = b"12345678" + try: + await mux.feed(build_ping(nonce).encode()) + finally: + await mux.close() + + assert len(sends) == 1 + assert sends[0][0] is True + frames = _decode_frames([sends[0][1]]) + assert len(frames) == 1 + assert frames[0].flags & FLAG_PONG + assert frames[0].payload == nonce + + @pytest.mark.asyncio async def test_wsgi_pool_survives_zero_credit_streams( tmp_path: Path, @@ -1265,9 +1353,17 @@ async def test_wsgi_pool_survives_zero_credit_streams( monkeypatch.setattr(wsgi_module, "WSGI_SEND_BRIDGE_POLL_SECONDS", 0.01) app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) _authorize_fingerprint(monkeypatch, fingerprint) + park_entries = 0 + park_entries_lock = threading.Lock() + park_entries_complete = threading.Event() @app.get("/_mux_test/park") def mux_test_park() -> Response: + nonlocal park_entries + with park_entries_lock: + park_entries += 1 + if park_entries == worker_count: + park_entries_complete.set() return Response( parked_body, content_type="application/octet-stream", @@ -1286,9 +1382,9 @@ async def test_wsgi_pool_survives_zero_credit_streams( healthy_closed = asyncio.Event() loop = asyncio.get_running_loop() identity = pl_identity(fingerprint) - executor = ThreadPoolExecutor(max_workers=worker_count) + admission = _admission(capacity=worker_count) - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: for frame in _decode_frames([data]): if frame.flags & FLAG_DATA: sent_payload_size[frame.stream_id] = sent_payload_size.get( @@ -1307,7 +1403,7 @@ async def test_wsgi_pool_survives_zero_credit_streams( healthy_closed.set() async def handler(reader: asyncio.StreamReader, writer: Any) -> None: - await dispatch_stream(app, identity, reader, writer, loop, executor) + await dispatch_stream(app, identity, reader, writer, loop, admission) async def open_get(stream_id: int, path: str) -> None: head = _http_head_bytes("GET", path, headers=None, content_length=0) @@ -1326,6 +1422,10 @@ async def test_wsgi_pool_survives_zero_credit_streams( for stream_id in parked_stream_ids: await open_get(stream_id, "/_mux_test/park") + assert await asyncio.wait_for( + asyncio.to_thread(park_entries_complete.wait), + timeout=1.0, + ) await asyncio.wait_for( asyncio.gather(*(event.wait() for event in zero_credit_events.values())), timeout=2.0, @@ -1337,7 +1437,7 @@ async def test_wsgi_pool_survives_zero_credit_streams( assert body == healthy_body finally: await mux.close() - await asyncio.to_thread(executor.shutdown, wait=True, cancel_futures=True) + await _shutdown_admission(admission) @pytest.mark.parametrize( @@ -1365,9 +1465,17 @@ async def test_teardown_paths_reclaim_wsgi_workers_for_sentinel_jobs( app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) fingerprint = "sha256:" + ("b" * 64) _authorize_fingerprint(monkeypatch, fingerprint) + park_entries = 0 + park_entries_lock = threading.Lock() + park_entries_complete = threading.Event() @app.get("/_mux_test/park") def mux_test_park() -> Response: + nonlocal park_entries + with park_entries_lock: + park_entries += 1 + if park_entries == worker_count: + park_entries_complete.set() return Response( b"x" * (INITIAL_WINDOW + 1), content_type="application/octet-stream", @@ -1378,9 +1486,9 @@ async def test_teardown_paths_reclaim_wsgi_workers_for_sentinel_jobs( writers: dict[int, Any] = {} loop = asyncio.get_running_loop() identity = pl_identity(fingerprint) - executor = ThreadPoolExecutor(max_workers=worker_count) + admission = _admission(capacity=worker_count) - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: for frame in _decode_frames([data]): if frame.flags & FLAG_DATA: sent_payload_size[frame.stream_id] = sent_payload_size.get( @@ -1393,7 +1501,7 @@ async def test_teardown_paths_reclaim_wsgi_workers_for_sentinel_jobs( async def handler(reader: asyncio.StreamReader, writer: Any) -> None: writers[writer.stream_id] = writer - await dispatch_stream(app, identity, reader, writer, loop, executor) + await dispatch_stream(app, identity, reader, writer, loop, admission) async def open_get(stream_id: int) -> None: head = _http_head_bytes( @@ -1409,6 +1517,10 @@ async def test_teardown_paths_reclaim_wsgi_workers_for_sentinel_jobs( for stream_id in stream_ids: await open_get(stream_id) + assert await asyncio.wait_for( + asyncio.to_thread(park_entries_complete.wait), + timeout=1.0, + ) await asyncio.wait_for( asyncio.gather(*(event.wait() for event in zero_credit_events.values())), timeout=2.0, @@ -1453,15 +1565,16 @@ async def test_teardown_paths_reclaim_wsgi_workers_for_sentinel_jobs( with contextlib.suppress(asyncio.CancelledError): await asyncio.wait_for(state.task, timeout=1.0) - sentinel_futures = [ - executor.submit(lambda value=value: value) for value in range(worker_count) - ] - assert [future.result(timeout=1.0) for future in sentinel_futures] == list( - range(worker_count) + sentinel_results = await asyncio.gather( + *( + admission.submit(loop, lambda value=value: value) + for value in range(worker_count) + ) ) + assert sentinel_results == list(range(worker_count)) finally: await mux.close() - await asyncio.to_thread(executor.shutdown, wait=True, cancel_futures=True) + await _shutdown_admission(admission) @pytest.mark.asyncio @@ -1469,7 +1582,7 @@ async def test_write_rechecks_closed_before_clearing_teardown_wake() -> None: emit_entered = asyncio.Event() release_emit = asyncio.Event() - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: emit_entered.set() await release_emit.wait() @@ -1510,7 +1623,7 @@ async def test_credit_stall_deadline_is_per_stall_and_window_resets_clock( data_events = [asyncio.Event() for _ in payload] write_task: asyncio.Task[None] | None = None - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: for frame in _decode_frames([data]): if frame.flags & FLAG_DATA: sent_bytes.append(frame.payload) @@ -1573,7 +1686,7 @@ async def test_send_credit_starvation_diagnostic_survives_closed_mux_without_sto captured: dict[str, Any] = {} state: Any | None = None - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: return async def handler(reader: asyncio.StreamReader, writer: Any) -> None: @@ -1636,15 +1749,15 @@ async def test_receive_window_credit_returns_after_head_and_body_consumption( loop = asyncio.get_running_loop() identity = pl_identity(fingerprint) - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) if any(frame.flags & FLAG_WINDOW for frame in _decode_frames([data])): window_event.set() - executor = ThreadPoolExecutor(max_workers=1) + admission = _admission(capacity=1) async def handler(reader: asyncio.StreamReader, writer: Any) -> None: - await dispatch_stream(app, identity, reader, writer, loop, executor) + await dispatch_stream(app, identity, reader, writer, loop, admission) mux = Multiplexer(send, handler, is_listener=True) try: @@ -1675,14 +1788,14 @@ async def test_receive_window_credit_returns_after_head_and_body_consumption( finally: await mux.close() assert mux._window_tasks == set() - await asyncio.to_thread(executor.shutdown, wait=True, cancel_futures=True) + await _shutdown_admission(admission) @pytest.mark.asyncio async def test_window_task_completion_logs_emit_failure( caplog: pytest.LogCaptureFixture, ) -> None: - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: return async def handler(*_: object) -> None: @@ -1731,13 +1844,13 @@ async def test_wsgi_input_read_timeout_releases_pool_worker( result: dict[str, int] = {} loop = asyncio.get_running_loop() identity = pl_identity(fingerprint) - executor = ThreadPoolExecutor(max_workers=1) + admission = _admission(capacity=1) - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(reader: asyncio.StreamReader, writer: Any) -> None: - dispatch = await dispatch_stream(app, identity, reader, writer, loop, executor) + dispatch = await dispatch_stream(app, identity, reader, writer, loop, admission) result["status"] = dispatch.status dispatch_done.set() @@ -1753,12 +1866,256 @@ async def test_wsgi_input_read_timeout_releases_pool_worker( await asyncio.wait_for(dispatch_done.wait(), timeout=1.0) assert result["status"] == 499 - assert ( - executor.submit(lambda: "worker-free").result(timeout=1.0) == "worker-free" - ) + assert await admission.submit(loop, lambda: "worker-free") == "worker-free" finally: await mux.close() - await asyncio.to_thread(executor.shutdown, wait=True, cancel_futures=True) + await _shutdown_admission(admission) + + +@pytest.mark.asyncio +async def test_capacity_snapshot_is_readable_while_workers_are_occupied( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) + admission = _admission(capacity=2, streaming_capacity=1) + entered = 0 + entered_lock = threading.Lock() + entered_all = threading.Event() + release_workers = threading.Event() + + def blocking_job() -> str: + nonlocal entered + with entered_lock: + entered += 1 + if entered == admission.config.capacity: + entered_all.set() + release_workers.wait(timeout=5.0) + return "done" + + tasks = [ + asyncio.create_task(admission.submit(asyncio.get_running_loop(), blocking_job)) + for _ in range(admission.config.capacity) + ] + try: + assert await asyncio.wait_for( + asyncio.to_thread(entered_all.wait), + timeout=1.0, + ) + + status, headers, body, writer = await _dispatch_raw_request( + app, + pl_identity("sha256:capacity"), + admission, + "GET", + wsgi_module.CAPACITY_SNAPSHOT_PATH, + ) + + payload = json.loads(body.decode("utf-8")) + assert status == 200 + assert headers["content-type"] == "application/json" + assert writer.closed is True + assert set(payload) == { + "timestamp_ms", + "active", + "limit", + "queued", + "rejected", + "admitted_over_budget", + "longest_wait_ms", + "refusal_enabled", + } + assert set(payload["active"]) == { + "total", + "streaming", + "streaming_over_budget", + } + assert set(payload["limit"]) == {"total", "streaming", "queue"} + assert set(payload["queued"]) == {"total"} + assert set(payload["rejected"]) == {"total", "streaming"} + assert set(payload["admitted_over_budget"]) == {"streaming"} + assert payload["active"]["total"] == admission.config.capacity + assert payload["limit"]["total"] == admission.config.capacity + assert "sha256" not in body.decode("utf-8") + assert "127.0.0.1" not in body.decode("utf-8") + finally: + release_workers.set() + await asyncio.gather(*tasks) + await _shutdown_admission(admission) + + +@pytest.mark.asyncio +async def test_capacity_snapshot_requires_paired_identity( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) + admission = _admission(capacity=1) + try: + status, _headers, _body, _writer = await _dispatch_raw_request( + app, + certless_identity(), + admission, + "GET", + wsgi_module.CAPACITY_SNAPSHOT_PATH, + ) + + assert status == 403 + finally: + await _shutdown_admission(admission) + + +@pytest.mark.asyncio +async def test_content_length_json_response_does_not_take_streaming_permit( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) + fingerprint = "sha256:" + ("e" * 64) + _authorize_fingerprint(monkeypatch, fingerprint) + + @app.get("/_mux_test/json") + def mux_test_json() -> Response: + return jsonify({"status": "ok"}) + + admission = _admission(capacity=1, streaming_capacity=1) + try: + status, headers, body, _writer = await _dispatch_raw_request( + app, + pl_identity(fingerprint), + admission, + "GET", + "/_mux_test/json", + ) + + snapshot = admission.snapshot() + assert status == 200 + assert headers["content-type"].startswith("application/json") + assert int(headers["content-length"]) == len(body) + assert snapshot["active"]["streaming"] == 0 + assert snapshot["rejected"]["streaming"] == 0 + assert snapshot["admitted_over_budget"]["streaming"] == 0 + finally: + await _shutdown_admission(admission) + + +@pytest.mark.asyncio +async def test_get_streaming_response_refuses_before_body_when_lane_full( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) + fingerprint = "sha256:" + ("f" * 64) + _authorize_fingerprint(monkeypatch, fingerprint) + + @app.get("/_mux_test/events") + def mux_test_events() -> Response: + def generate() -> Iterator[bytes]: + yield b"data: should-not-send\n\n" + + return Response(generate(), mimetype="text/event-stream") + + admission = _admission(capacity=1, streaming_capacity=1) + held = admission.try_acquire_streaming() + assert held is not None + try: + status, _headers, body, writer = await _dispatch_raw_request( + app, + pl_identity(fingerprint), + admission, + "GET", + "/_mux_test/events", + ) + + snapshot = admission.snapshot() + assert status == 503 + assert writer.closed is True + assert b"secure listener streaming capacity is full" in body + assert b"should-not-send" not in body + assert snapshot["rejected"]["streaming"] == 1 + assert snapshot["admitted_over_budget"]["streaming"] == 0 + finally: + held.release() + await _shutdown_admission(admission) + + +@pytest.mark.asyncio +async def test_post_streaming_response_proceeds_over_budget_when_lane_full( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(wsgi_module, "STREAMING_PERMIT_WAIT_TIMEOUT_SECONDS", 0.0) + app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) + fingerprint = "sha256:" + ("0" * 64) + _authorize_fingerprint(monkeypatch, fingerprint) + + @app.post("/_mux_test/events") + def mux_test_events() -> Response: + def generate() -> Iterator[bytes]: + yield b"data: committed\n\n" + + return Response(generate(), mimetype="text/event-stream") + + admission = _admission(capacity=1, streaming_capacity=1) + held = admission.try_acquire_streaming() + assert held is not None + try: + status, _headers, body, writer = await _dispatch_raw_request( + app, + pl_identity(fingerprint), + admission, + "POST", + "/_mux_test/events", + ) + + snapshot = admission.snapshot() + assert status == 200 + assert writer.closed is True + assert body == b"data: committed\n\n" + assert snapshot["rejected"]["streaming"] == 0 + assert snapshot["active"]["streaming_over_budget"] == 0 + assert snapshot["admitted_over_budget"]["streaming"] == 1 + finally: + held.release() + await _shutdown_admission(admission) + + +@pytest.mark.asyncio +async def test_streaming_capacity_zero_disables_lane_split( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) + fingerprint = "sha256:" + ("1" * 64) + _authorize_fingerprint(monkeypatch, fingerprint) + + @app.get("/_mux_test/events") + def mux_test_events() -> Response: + def generate() -> Iterator[bytes]: + yield b"data: legacy-shared-pool\n\n" + + return Response(generate(), mimetype="text/event-stream") + + admission = _admission(capacity=1, streaming_capacity=0) + try: + status, _headers, body, writer = await _dispatch_raw_request( + app, + pl_identity(fingerprint), + admission, + "GET", + "/_mux_test/events", + ) + + snapshot = admission.snapshot() + assert status == 200 + assert writer.closed is True + assert body == b"data: legacy-shared-pool\n\n" + assert snapshot["limit"]["streaming"] == 0 + assert snapshot["active"]["streaming"] == 0 + assert snapshot["rejected"]["streaming"] == 0 + assert snapshot["admitted_over_budget"]["streaming"] == 0 + finally: + await _shutdown_admission(admission) @pytest.mark.asyncio @@ -1771,17 +2128,18 @@ async def test_early_bridge_response_drains_in_flight_body_without_unknown_strea sent: list[bytes] = [] loop = asyncio.get_running_loop() - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) - with ThreadPoolExecutor(max_workers=1) as executor: + admission = _admission(capacity=1) + try: async def handler( reader: asyncio.StreamReader, writer: Any, ) -> None: stream_identity = certless_identity() if writer.stream_id == 3 else identity - await dispatch_stream(app, stream_identity, reader, writer, loop, executor) + await dispatch_stream(app, stream_identity, reader, writer, loop, admission) mux = Multiplexer(send, handler, is_listener=True) try: @@ -1806,6 +2164,8 @@ async def test_early_bridge_response_drains_in_flight_body_without_unknown_strea assert state.task is not None and state.task.done() finally: await mux.close() + finally: + await _shutdown_admission(admission) @pytest.mark.asyncio @@ -1819,17 +2179,18 @@ async def test_early_bridge_response_cancels_on_drain_budget_exhaustion( diags: list[ResetDiagnostic] = [] loop = asyncio.get_running_loop() - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) - with ThreadPoolExecutor(max_workers=1) as executor: + admission = _admission(capacity=1) + try: async def handler( reader: asyncio.StreamReader, writer: Any, ) -> None: stream_identity = certless_identity() if writer.stream_id == 3 else identity - await dispatch_stream(app, stream_identity, reader, writer, loop, executor) + await dispatch_stream(app, stream_identity, reader, writer, loop, admission) mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) try: @@ -1861,6 +2222,8 @@ async def test_early_bridge_response_cancels_on_drain_budget_exhaustion( ) finally: await mux.close() + finally: + await _shutdown_admission(admission) @pytest.mark.asyncio @@ -1874,17 +2237,18 @@ async def test_app_early_response_drains_then_cancels_with_app_context( diags: list[ResetDiagnostic] = [] loop = asyncio.get_running_loop() - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) - with ThreadPoolExecutor(max_workers=1) as executor: + admission = _admission(capacity=1) + try: async def handler( reader: asyncio.StreamReader, writer: Any, ) -> None: stream_identity = certless_identity() if writer.stream_id == 3 else identity - await dispatch_stream(app, stream_identity, reader, writer, loop, executor) + await dispatch_stream(app, stream_identity, reader, writer, loop, admission) mux = Multiplexer(send, handler, is_listener=True, on_reset=diags.append) try: @@ -1924,6 +2288,8 @@ async def test_app_early_response_drains_then_cancels_with_app_context( assert not mux._closed finally: await mux.close() + finally: + await _shutdown_admission(admission) @pytest.mark.asyncio @@ -1931,7 +2297,7 @@ async def test_handler_exception_is_stream_scoped() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler( @@ -1974,7 +2340,7 @@ async def test_true_mux_violations_keep_reason_codes() -> None: sent: list[bytes] = [] diags: list[ResetDiagnostic] = [] - async def send(data: bytes) -> None: + async def send(data: bytes, *, urgent: bool = False) -> None: sent.append(data) async def handler(*_: object) -> None: @@ -2071,7 +2437,7 @@ async def test_reset_diagnostics_distinguish_contexts_and_are_privacy_clean( ) -> None: diags: list[ResetDiagnostic] = [] - async def send(_: bytes) -> None: + async def send(_: bytes, *, urgent: bool = False) -> None: return async def collect_with_handler(handler: Any, frame: bytes) -> None: @@ -2163,7 +2529,8 @@ async def test_reset_diagnostics_distinguish_contexts_and_are_privacy_clean( app, _journal = make_convey_app(tmp_path, monkeypatch, link={"posture": "spl"}) loop = asyncio.get_running_loop() - with ThreadPoolExecutor(max_workers=1) as executor: + admission = _admission(capacity=1) + try: async def bridge_handler( reader: asyncio.StreamReader, @@ -2175,7 +2542,7 @@ async def test_reset_diagnostics_distinguish_contexts_and_are_privacy_clean( reader, writer, loop, - executor, + admission, ) bridge_mux = Multiplexer( @@ -2207,7 +2574,7 @@ async def test_reset_diagnostics_distinguish_contexts_and_are_privacy_clean( reader, writer, loop, - executor, + admission, ) app_mux = Multiplexer( @@ -2225,6 +2592,8 @@ async def test_reset_diagnostics_distinguish_contexts_and_are_privacy_clean( await app_mux.feed(build_data(1, b"x" * state.recv_credit).encode()) finally: await app_mux.close() + finally: + await _shutdown_admission(admission) contexts = {diag.context for diag in diags} assert { diff --git a/tests/link/test_wsgi_location_rewrite.py b/tests/link/test_wsgi_location_rewrite.py index c3064dd10..ac8d285fc 100644 --- a/tests/link/test_wsgi_location_rewrite.py +++ b/tests/link/test_wsgi_location_rewrite.py @@ -4,13 +4,16 @@ from __future__ import annotations import asyncio -from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import Any import pytest from solstone.convey import root as root_module +from solstone.convey.secure_listener.admission import ( + SecureListenerAdmission, + SecureListenerAdmissionConfig, +) from solstone.convey.secure_listener.identity import ConveyIdentity from solstone.convey.secure_listener.wsgi import ( _normalize_location_headers, @@ -56,8 +59,17 @@ async def _dispatch_raw_request( reader.feed_eof() writer = FakeStreamWriter() loop = asyncio.get_running_loop() - with ThreadPoolExecutor(max_workers=1) as executor: - await dispatch_stream(app, identity, reader, writer, loop, executor) + admission = SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=1, + streaming_capacity=1, + refuse_when_full=False, + ) + ) + try: + await dispatch_stream(app, identity, reader, writer, loop, admission) + finally: + await asyncio.to_thread(admission.shutdown, wait=True, cancel_futures=True) status, headers, body = _parse_http_response(bytes(writer.data)) return status, headers, body, writer diff --git a/tests/test_secure_listener_runtime.py b/tests/test_secure_listener_runtime.py index da8d605d1..7e1c2515e 100644 --- a/tests/test_secure_listener_runtime.py +++ b/tests/test_secure_listener_runtime.py @@ -7,13 +7,13 @@ import json import logging import socket import threading -from concurrent.futures import ThreadPoolExecutor from pathlib import Path from types import SimpleNamespace from unittest.mock import MagicMock import pytest +from solstone.convey.secure_listener import accept as accept_module from solstone.convey.secure_listener import runtime as rt from solstone.convey.secure_listener.accept import ( CERTLESS_PAIR_FAILURE_CAP, @@ -22,6 +22,15 @@ from solstone.convey.secure_listener.accept import ( SecureListener, certless_admission_mode, ) +from solstone.convey.secure_listener.admission import ( + DEFAULT_SECURE_LISTENER_CAPACITY, + DEFAULT_SECURE_LISTENER_STREAMING_CAPACITY, + SecureListenerAdmission, + SecureListenerAdmissionConfig, + SecureListenerAdmissionRejected, + resolve_admission_config, +) +from solstone.convey.secure_listener.framing import build_ping from solstone.think.link import client as link_client from solstone.think.link.ca import load_or_generate_ca from solstone.think.link.nonces import NONCE_TTL_SECONDS, NonceStore @@ -31,13 +40,19 @@ from tests.link.secure_listener_harness import SecureListenerHarness def test_reuse_port_allows_coexisting_bind(): - executor = ThreadPoolExecutor(max_workers=1) + admission = SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=1, + streaming_capacity=1, + refuse_when_full=False, + ) + ) listener = SecureListener( app=MagicMock(), strict_tls_ctx=MagicMock(), relaxed_tls_ctx=MagicMock(), authorized=set(), - executor=executor, + admission=admission, callosum_emit=lambda *a, **kw: None, host="127.0.0.1", port=0, @@ -60,7 +75,7 @@ def test_reuse_port_allows_coexisting_bind(): if s2 is not None: s2.close() loop.close() - executor.shutdown(wait=True, cancel_futures=True) + admission.shutdown(wait=True, cancel_futures=True) def test_stop_all_after_loop_closed_does_not_raise(): @@ -68,7 +83,13 @@ def test_stop_all_after_loop_closed_does_not_raise(): previous_runtime = rt._runtime s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - executor = ThreadPoolExecutor(max_workers=1) + admission = SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=1, + streaming_capacity=1, + refuse_when_full=False, + ) + ) try: s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEPORT, 1) @@ -86,7 +107,7 @@ def test_stop_all_after_loop_closed_does_not_raise(): loop=loop, thread=thread, apps=[app], - executor=executor, + admission=admission, listener=listener, sockets=(s,), ) @@ -98,7 +119,114 @@ def test_stop_all_after_loop_closed_does_not_raise(): finally: rt._runtime = previous_runtime s.close() - executor.shutdown(wait=True, cancel_futures=True) + admission.shutdown(wait=True, cancel_futures=True) + + +def test_secure_listener_admission_config_defaults_to_current_capacity( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _runtime_journal( + tmp_path, + monkeypatch, + {"setup": {"completed_at": 1700000000000}}, + ) + + config = resolve_admission_config() + + assert config.capacity == DEFAULT_SECURE_LISTENER_CAPACITY == 16 + assert config.streaming_capacity == DEFAULT_SECURE_LISTENER_STREAMING_CAPACITY == 8 + assert config.refuse_when_full is False + assert config.queue_limit == 32 + + +def test_secure_listener_admission_config_reads_link_namespace( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _runtime_journal( + tmp_path, + monkeypatch, + { + "setup": {"completed_at": 1700000000000}, + "link": { + "secure_listener_capacity": 24, + "secure_listener_streaming_capacity": 6, + "secure_listener_refuse_when_full": True, + }, + }, + ) + + config = resolve_admission_config() + + assert config.capacity == 24 + assert config.streaming_capacity == 6 + assert config.refuse_when_full is True + assert config.queue_limit == 48 + + +def test_secure_listener_admission_config_warns_and_falls_back( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + _runtime_journal( + tmp_path, + monkeypatch, + { + "setup": {"completed_at": 1700000000000}, + "link": { + "secure_listener_capacity": "wide", + "secure_listener_streaming_capacity": "many", + "secure_listener_refuse_when_full": "yes", + }, + }, + ) + + with caplog.at_level( + logging.WARNING, + logger="convey.secure_listener.admission", + ): + config = resolve_admission_config() + + assert config == SecureListenerAdmissionConfig() + assert ( + "Invalid link.secure_listener_capacity in journal config: 'wide' " + "\u2014 defaulting to 16" + ) in caplog.text + assert ( + "Invalid link.secure_listener_streaming_capacity in journal config: 'many' " + "\u2014 defaulting to 8" + ) in caplog.text + assert ( + "Invalid link.secure_listener_refuse_when_full in journal config: 'yes' " + "\u2014 defaulting to false" + ) in caplog.text + + +@pytest.mark.asyncio +async def test_secure_listener_admission_refuses_when_enabled_queue_is_full() -> None: + admission = SecureListenerAdmission( + SecureListenerAdmissionConfig( + capacity=1, + streaming_capacity=0, + refuse_when_full=True, + ) + ) + try: + with admission._lock: + admission._active_total = admission.config.capacity + for _ in range(admission.config.queue_limit): + admission._waiters.append(SimpleNamespace(queued_at=0.0)) + + with pytest.raises(SecureListenerAdmissionRejected): + await admission.acquire() + + snapshot = admission.snapshot() + assert snapshot["queued"]["total"] == admission.config.queue_limit + assert snapshot["rejected"]["total"] == 1 + finally: + admission.shutdown(wait=True, cancel_futures=True) def test_start_secure_listener_setup_incomplete_does_not_establish_identity( @@ -263,6 +391,128 @@ def test_certless_admission_mode( assert certless_admission_mode(mode, window_is_open) == expected +@pytest.mark.asyncio +async def test_priority_send_queue_drains_urgent_frames_first_and_preserves_join( + monkeypatch: pytest.MonkeyPatch, +) -> None: + queue: asyncio.PriorityQueue[tuple[int, int, bytes]] = asyncio.PriorityQueue() + queue.put_nowait((1, 0, b"normal")) + queue.put_nowait((0, 1, b"urgent")) + written: list[bytes] = [] + + monkeypatch.setattr(accept_module, "_encrypt", lambda _tls, plaintext: plaintext) + + async def write_ciphertext(data: bytes) -> None: + written.append(data) + + await accept_module._drain_send_queue(object(), write_ciphertext, queue) + await asyncio.wait_for(queue.join(), timeout=1.0) + + assert written == [b"urgent", b"normal"] + + +def test_tcp_keepalive_options_are_feature_detected() -> None: + linux_socket = SimpleNamespace( + SOL_SOCKET=1, + SO_KEEPALIVE=2, + IPPROTO_TCP=3, + TCP_KEEPIDLE=4, + TCP_KEEPINTVL=5, + TCP_KEEPCNT=6, + ) + mac_socket = SimpleNamespace( + SOL_SOCKET=1, + SO_KEEPALIVE=2, + IPPROTO_TCP=3, + TCP_KEEPALIVE=7, + ) + minimal_socket = SimpleNamespace(SOL_SOCKET=1, SO_KEEPALIVE=2) + + assert accept_module._tcp_keepalive_options(linux_socket) == [ + (1, 2, 1), + (3, 4, 30), + (3, 5, 10), + (3, 6, 3), + ] + assert accept_module._tcp_keepalive_options(mac_socket) == [ + (1, 2, 1), + (3, 7, 30), + ] + assert accept_module._tcp_keepalive_options(minimal_socket) == [(1, 2, 1)] + + +def test_tcp_keepalive_failures_log_and_continue( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + sock = _FailingSocket() + writer = _SocketWriter(sock) + monkeypatch.setattr( + accept_module, + "_tcp_keepalive_options", + lambda: [(1, 2, 3), (4, 5, 6)], + ) + + with caplog.at_level(logging.WARNING, logger="convey.secure_listener.accept"): + accept_module._apply_tcp_keepalive( + writer, logging.getLogger("convey.secure_listener.accept") + ) + + assert sock.calls == [(1, 2, 3), (4, 5, 6)] + assert caplog.text.count("secure listener TCP keepalive option failed") == 2 + + +@pytest.mark.asyncio +async def test_pump_connection_writer_failure_ends_connection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + listener = _listener() + tcp_reader = asyncio.StreamReader() + tcp_reader.feed_data(b"tls-record") + tcp_writer = _HangingDrainWriter() + ping = build_ping(b"12345678").encode() + tls = SimpleNamespace(handshake_done=False, peer_fingerprint=None) + + monkeypatch.setattr(accept_module, "new_server", lambda _ctx: tls) + monkeypatch.setattr(accept_module, "window_open", lambda: False) + monkeypatch.setattr( + accept_module, + "SECURE_LISTENER_TCP_DRAIN_TIMEOUT_SECONDS", + 0.0, + ) + + async def no_reader_drain(*_args: object) -> None: + return + + def fake_drive_tls( + _tls: object, + *, + inbound: bytes = b"", + plaintext_out: bytes = b"", + ) -> tuple[bytes, bytes]: + if plaintext_out: + return plaintext_out, b"" + if inbound: + return b"", ping + return b"", b"" + + monkeypatch.setattr(accept_module, "_drain_send_queue", no_reader_drain) + monkeypatch.setattr(accept_module, "drive_tls", fake_drive_tls) + + with pytest.raises(TimeoutError): + await asyncio.wait_for( + listener._pump_connection( + tcp_reader, + tcp_writer, + "conn-writer-failure", + "pl-direct", + ), + timeout=1.0, + ) + + assert tcp_writer.writes + + @pytest.mark.asyncio async def test_certless_reap_tears_down_on_passive_expiry( tmp_path: Path, @@ -481,7 +731,7 @@ def _listener() -> SecureListener: strict_tls_ctx=MagicMock(), relaxed_tls_ctx=MagicMock(), authorized=MagicMock(), - executor=MagicMock(), + admission=MagicMock(), callosum_emit=lambda *a, **kw: None, host="127.0.0.1", port=0, @@ -553,3 +803,33 @@ class _FakeMux: async def close(self) -> None: self.closed = True + + +class _FailingSocket: + def __init__(self) -> None: + self.calls: list[tuple[int, int, int]] = [] + + def setsockopt(self, level: int, optname: int, value: int) -> None: + self.calls.append((level, optname, value)) + raise OSError("option unavailable") + + +class _SocketWriter: + def __init__(self, sock: _FailingSocket) -> None: + self._sock = sock + + def get_extra_info(self, name: str) -> object: + if name == "socket": + return self._sock + return None + + +class _HangingDrainWriter: + def __init__(self) -> None: + self.writes: list[bytes] = [] + + def write(self, data: bytes) -> None: + self.writes.append(data) + + async def drain(self) -> None: + await asyncio.Event().wait()