diff --git a/benchmarks/route_benchmark.py b/benchmarks/route_benchmark.py new file mode 100644 index 0000000..1efab47 --- /dev/null +++ b/benchmarks/route_benchmark.py @@ -0,0 +1,125 @@ +"""Repeatable routing comparison; run with ``python benchmarks/route_benchmark.py``.""" + +from __future__ import annotations + +import argparse +import asyncio +import importlib.util +import inspect +import json +from pathlib import Path +import statistics +import sys +import time + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from smallserver import Headers, RegexRouteConfig, Request, Response, RouteMatchTimeout, SmallServer + +STATIC_DISPATCH_RATIO_FLOOR = 0.80 + + +async def _legacy_dispatch(routes, request): + """Model the pre-router static dictionary dispatch for same-run comparison.""" + handler = routes[(request.method.upper(), request.path)] + result = handler(request) + if not inspect.isawaitable(result): + raise TypeError("benchmark handler must be awaitable") + response = await result + if not isinstance(response, Response): + raise TypeError("benchmark handler must return Response") + return response + + +async def _measure(operation, iterations: int) -> float: + started = time.perf_counter() + for _ in range(iterations): + await operation() + return iterations / (time.perf_counter() - started) + + +async def benchmark(iterations: int, rounds: int) -> dict[str, float | int | str | bool]: + app = SmallServer() + + @app.get("/health") + async def health(request): + return Response() + + request = Request("GET", "/health", Headers()) + legacy_routes = {("GET", "/health"): health} + + async def legacy_operation(): + return await _legacy_dispatch(legacy_routes, request) + + async def router_operation(): + return await app.dispatch(request) + + await _measure(legacy_operation, min(iterations, 1_000)) + await _measure(router_operation, min(iterations, 1_000)) + legacy_rates = [] + router_rates = [] + ratios = [] + for round_number in range(rounds): + if round_number % 2: + router_rate = await _measure(router_operation, iterations) + legacy_rate = await _measure(legacy_operation, iterations) + else: + legacy_rate = await _measure(legacy_operation, iterations) + router_rate = await _measure(router_operation, iterations) + legacy_rates.append(legacy_rate) + router_rates.append(router_rate) + ratios.append(router_rate / legacy_rate) + + median_ratio = statistics.median(ratios) + result: dict[str, float | int | str | bool] = { + "iterations": iterations, + "rounds": rounds, + "legacy_static_dispatches_per_second": round(statistics.median(legacy_rates), 2), + "router_static_dispatches_per_second": round(statistics.median(router_rates), 2), + "router_to_legacy_ratio": round(median_ratio, 4), + "static_dispatch_ratio_floor": STATIC_DISPATCH_RATIO_FLOOR, + "static_dispatch_floor_passed": median_ratio >= STATIC_DISPATCH_RATIO_FLOOR, + } + + if importlib.util.find_spec("regex") is None: + result["regex"] = "skipped; install smallserver[regex-routes]" + return result + + bounded = SmallServer(RegexRouteConfig(match_timeout=0.002, total_match_timeout=0.005)) + + @bounded.get_regex(r"/(a+)+$") + async def hostile(request): + return Response() + + started = time.perf_counter() + try: + await bounded.dispatch(Request("GET", "/" + "a" * 5000 + "!", Headers())) + except RouteMatchTimeout: + pass + result["configured_regex_match_timeout_seconds"] = 0.002 + result["observed_worst_case_regex_seconds"] = round(time.perf_counter() - started, 6) + return result + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--iterations", type=int, default=25_000) + parser.add_argument("--rounds", type=int, default=5) + parser.add_argument("--release", action="store_true") + arguments = parser.parse_args() + if arguments.iterations <= 0: + parser.error("--iterations must be positive") + if arguments.rounds <= 0: + parser.error("--rounds must be positive") + if arguments.release and arguments.iterations < 10_000: + parser.error("--release requires at least 10000 iterations") + if arguments.release and arguments.rounds < 5: + parser.error("--release requires at least 5 rounds") + result = asyncio.run(benchmark(arguments.iterations, arguments.rounds)) + print(json.dumps(result, indent=2, sort_keys=True)) + if arguments.release and not result["static_dispatch_floor_passed"]: + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/demo.py b/demo.py index bfa2f07..1fc7613 100644 --- a/demo.py +++ b/demo.py @@ -4,7 +4,7 @@ import json -from smallserver import HTTPError, Request, Response, SmallServer +from smallserver import HTTPError, Request, Response, SmallServer, WebSocket app = SmallServer() @@ -81,6 +81,16 @@ async def delete_task(request: Request) -> Response: return Response(status=204) +@app.websocket("/ws") +async def websocket_echo(socket: WebSocket) -> None: + await socket.accept() + async for message in socket: + if message.is_text: + await socket.send_text(message.text) + else: + await socket.send_bytes(message.bytes) + + if __name__ == "__main__": print("Starting SmallServer on http://127.0.0.1:8000") app.listen(host="127.0.0.1", port=8000) diff --git a/examples/websocket_echo.py b/examples/websocket_echo.py new file mode 100644 index 0000000..4153dea --- /dev/null +++ b/examples/websocket_echo.py @@ -0,0 +1,28 @@ +"""Run a bounded WebSocket echo endpoint on localhost:8000.""" + +from smallserver import SmallServer, WebSocket, WebSocketConfig + + +app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=64 * 1024, + max_message_bytes=256 * 1024, + max_inbound_messages=8, + max_outbound_commands=8, + ) +) + + +@app.websocket("/echo") +async def echo(socket: WebSocket) -> None: + await socket.accept() + async for message in socket: + if message.is_text: + await socket.send_text(message.text) + else: + await socket.send_bytes(message.bytes) + + +if __name__ == "__main__": + print("Starting WebSocket echo server on ws://127.0.0.1:8000/echo") + app.listen(host="127.0.0.1", port=8000) diff --git a/pyproject.toml b/pyproject.toml index 800e3cd..30ab6ee 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,6 +14,8 @@ dependencies = [] [project.optional-dependencies] dev = ["build>=1.2"] +regex-routes = ["regex>=2023.10.3,<2027"] +websocket = ["wsproto>=1.2,<2"] test = [ "build>=1.2", "h2>=4,<5", diff --git a/smallserver/__init__.py b/smallserver/__init__.py index a3f0a16..b297939 100644 --- a/smallserver/__init__.py +++ b/smallserver/__init__.py @@ -10,8 +10,24 @@ ServerStartupError, ) from .http import Headers, Request, Response +from .routing import ( + RegexRouteConfig, + RegexRoutesUnavailable, + RouteErrorEvent, + RouteMatchTimeout, + RoutePathTooLarge, +) from .runtime import ManagedRuntimeConfig from .server import ServerConfig, ServerHandle +from .websocket import ( + WebSocket, + WebSocketCapacityError, + WebSocketConfig, + WebSocketDisconnect, + WebSocketMessage, + WebSocketStateError, + WebSocketUnavailable, +) if TYPE_CHECKING: from .adapters import AdapterRegistry, AdapterShutdownError, http_error_from_adapter @@ -33,13 +49,25 @@ def __getattr__(name: str) -> Any: "Headers", "HTTPError", "ManagedRuntimeConfig", + "RegexRouteConfig", + "RegexRoutesUnavailable", "Request", "Response", + "RouteErrorEvent", + "RouteMatchTimeout", + "RoutePathTooLarge", "ServerConfig", "ServerConfigurationError", "ServerFinalizationError", "ServerHandle", "ServerStartupError", "SmallServer", + "WebSocket", + "WebSocketCapacityError", + "WebSocketConfig", + "WebSocketDisconnect", + "WebSocketMessage", + "WebSocketStateError", + "WebSocketUnavailable", "http_error_from_adapter", ] diff --git a/smallserver/app.py b/smallserver/app.py index 08343ff..68ffb1f 100644 --- a/smallserver/app.py +++ b/smallserver/app.py @@ -4,6 +4,7 @@ import inspect from collections.abc import Awaitable, Callable, Iterable +from dataclasses import replace from typing import Any, Literal, NoReturn, Protocol, cast, overload try: @@ -24,11 +25,37 @@ _CleanupTransaction, ) from .http import Request, Response +from .routing import ( + RegexRouteConfig, + RouteErrorEvent, + RouteMatchTimeout, + RoutePathTooLarge, + Router, +) from .runtime import ManagedRuntimeConfig -from .server import HTTPParseError, HTTPRequestParser, ServerConfig, ServerHandle +from .server import ( + HTTPParseError, + HTTPRequestParser, + RouteObserverChannel, + ServerConfig, + ServerHandle, + run_route_observer, +) +from .websocket import ( + WebSocket, + WebSocketConfig, + WebSocketUnavailable, + _WebSocketRoute, + _WebSocketState, + _is_http_token, + _is_upgrade_attempt, + _validate_upgrade, + run_websocket_connection, +) Handler = Callable[[Request], Awaitable[Response]] -_METHODS = frozenset({"GET", "POST", "PUT", "PATCH", "DELETE"}) +WebSocketHandler = Callable[[WebSocket], Awaitable[None]] +RouteErrorObserver = Callable[[RouteErrorEvent], None] class _NoThreadLock: @@ -123,8 +150,24 @@ def errors(self) -> tuple[BaseException, ...]: class SmallServer: """Register static HTTP routes and dispatch requests to async handlers.""" - def __init__(self) -> None: - self._routes: dict[tuple[str, str], Handler] = {} + def __init__( + self, + regex_config: RegexRouteConfig | None = None, + *, + route_error_observer: RouteErrorObserver | None = None, + websocket_config: WebSocketConfig | None = None, + ) -> None: + if route_error_observer is not None and not callable(route_error_observer): + raise TypeError("route_error_observer must be callable or None") + self._router = Router(regex_config) + self._routes = self._router._static + self._route_error_observer = route_error_observer + if websocket_config is not None and not isinstance( + websocket_config, WebSocketConfig + ): + raise TypeError("websocket_config must be a WebSocketConfig or None") + self._websocket_config = websocket_config or WebSocketConfig() + self._websocket_routes: dict[str, _WebSocketRoute] = {} self._active_invocation: object | ServerHandle | None = None self._invocation_lock: Any = ( allocate_lock() if allocate_lock is not None else _NoThreadLock() @@ -154,20 +197,14 @@ def _release_invocation(self, expected: object) -> None: def route(self, path: str, methods: Iterable[str]) -> Callable[[Handler], Handler]: if not isinstance(path, str) or not path.startswith("/"): raise ValueError("route path must start with '/'") - normalized = tuple(dict.fromkeys(method.upper() for method in methods)) - if not normalized or any(method not in _METHODS for method in normalized): - raise ValueError("routes must use one or more supported HTTP methods") + if "?" in path or "#" in path: + raise ValueError("route path must not contain a query string or fragment") + normalized = self._router.normalize_methods(methods) def register(handler: Handler) -> Handler: if not callable(handler): raise TypeError("route handler must be callable") - keys = [(method, path) for method in normalized] - for method, key_path in keys: - key = (method, key_path) - if key in self._routes: - raise ValueError("route already registered: {} {}".format(method, path)) - for key in keys: - self._routes[key] = handler + self._router.add_static(path, normalized, handler) return handler return register @@ -187,6 +224,75 @@ def patch(self, path: str) -> Callable[[Handler], Handler]: def delete(self, path: str) -> Callable[[Handler], Handler]: return self.route(path, ("DELETE",)) + def websocket( + self, + path: str, + *, + origins: Iterable[str] | None = None, + subprotocols: Iterable[str] = (), + ) -> Callable[[WebSocketHandler], WebSocketHandler]: + """Register a static HTTP/1.1 WebSocket Upgrade route.""" + if not isinstance(path, str) or not path.startswith("/"): + raise ValueError("WebSocket route path must start with '/'") + if "?" in path or "#" in path: + raise ValueError("WebSocket route path must not contain query or fragment") + if path in self._websocket_routes: + raise ValueError("WebSocket route already registered: {}".format(path)) + origin_set: frozenset[str] | None = None + if origins is not None: + if isinstance(origins, str): + raise TypeError("WebSocket origins must be an iterable of strings") + origin_set = frozenset(origins) + if any(not isinstance(origin, str) or not origin for origin in origin_set): + raise ValueError("WebSocket origins must be non-empty strings") + if isinstance(subprotocols, str): + raise TypeError("WebSocket subprotocols must be an iterable of tokens") + protocols = tuple(dict.fromkeys(subprotocols)) + if any( + not isinstance(protocol, str) or not _is_http_token(protocol) + for protocol in protocols + ): + raise ValueError("WebSocket subprotocols must be valid HTTP tokens") + + def register(handler: WebSocketHandler) -> WebSocketHandler: + if not callable(handler): + raise TypeError("WebSocket handler must be callable") + if path in self._websocket_routes: + raise ValueError("WebSocket route already registered: {}".format(path)) + self._websocket_routes[path] = _WebSocketRoute( + handler, origin_set, protocols + ) + return handler + + return register + + def route_regex(self, pattern: str, methods: Iterable[str]) -> Callable[[Handler], Handler]: + """Register a timeout-bounded full-path regular-expression route.""" + normalized = self._router.normalize_methods(methods) + + def register(handler: Handler) -> Handler: + if not callable(handler): + raise TypeError("route handler must be callable") + self._router.add_regex(pattern, normalized, handler) + return handler + + return register + + def get_regex(self, pattern: str) -> Callable[[Handler], Handler]: + return self.route_regex(pattern, ("GET",)) + + def post_regex(self, pattern: str) -> Callable[[Handler], Handler]: + return self.route_regex(pattern, ("POST",)) + + def put_regex(self, pattern: str) -> Callable[[Handler], Handler]: + return self.route_regex(pattern, ("PUT",)) + + def patch_regex(self, pattern: str) -> Callable[[Handler], Handler]: + return self.route_regex(pattern, ("PATCH",)) + + def delete_regex(self, pattern: str) -> Callable[[Handler], Handler]: + return self.route_regex(pattern, ("DELETE",)) + def serve( self, runtime: _RuntimeLike, @@ -290,9 +396,8 @@ def listen( def _handle_cleanup_transaction(handle: ServerHandle) -> _CleanupTransaction: return _HandleCleanupTransaction(handle) - @staticmethod def _resolve_server_config( - config: ServerConfig | None, *, managed: bool + self, config: ServerConfig | None, *, managed: bool ) -> ServerConfig: if config is not None and not isinstance(config, ServerConfig): raise TypeError("config must be a ServerConfig or None") @@ -313,11 +418,14 @@ def _resolve_server_config( "server task priorities must be lower than managed runtime " "priority_levels" ) - required_tasks = resolved.max_connections + 2 + control_tasks = 3 if self._route_error_observer is not None else 2 + required_tasks = resolved.max_connections + control_tasks if effective_runtime_config.task_capacity < required_tasks: raise ValueError( "managed runtime task_capacity must be at least " - "max_connections + 2 for listener and shutdown tasks" + "max_connections + {} for server control tasks".format( + control_tasks + ) ) return resolved @@ -400,8 +508,21 @@ def release(completed: ServerHandle) -> None: self._release_invocation(completed) try: + observer_channel = ( + RouteObserverChannel( + self._route_error_observer, config.max_route_error_events + ) + if self._route_error_observer is not None + else None + ) handle = ServerHandle( - runtime, transport, listener, wakeup, config, on_finalized=release + runtime, + transport, + listener, + wakeup, + config, + on_finalized=release, + route_observer_channel=observer_channel, ) except BaseException as primary_error: transaction = _CleanupTransaction() @@ -442,6 +563,16 @@ def release(completed: ServerHandle) -> None: tasks = (listener_task, close_task) handle._close_task = close_task handle._owned_tasks.append(close_task) + if observer_channel is not None: + observer_task = SmallTask( + config.listener_priority, + run_route_observer, + args=(observer_channel,), + name="smallserver-route-observer", + ) + observer_channel.bind(observer_task) + tasks += (observer_task,) + handle._owned_tasks.append(observer_task) runtime.fork(list(tasks)) except BaseException as primary_error: handle._abort_startup(tasks) @@ -453,12 +584,21 @@ def release(completed: ServerHandle) -> None: async def dispatch(self, request: Request) -> Response: """Run a registered handler or return a deterministic HTTP response.""" - handler = self._routes.get((request.method.upper(), request.path)) - if handler is None: - allowed = sorted(method for method, path in self._routes if path == request.path) - if allowed: - return Response.text("method not allowed", status=405, headers={"Allow": ", ".join(allowed)}) - return Response.text("not found", status=404) + handler = self._router.static_handler(request.method, request.path) + if handler is not None: + if request.path_params or request.route_pattern is not None: + request = replace(request, path_params={}, route_pattern=None) + else: + try: + match = self._router.resolve(request.method, request.path) + except RoutePathTooLarge: + return Response.text("request target is too large", status=414) + if match.handler is None: + if match.allowed_methods: + return Response.text("method not allowed", status=405, headers={"Allow": ", ".join(match.allowed_methods)}) + return Response.text("not found", status=404) + handler = match.handler + request = replace(request, path_params=match.path_params, route_pattern=match.route_pattern) try: result = handler(request) if not inspect.isawaitable(result): @@ -544,8 +684,11 @@ async def _connection_loop( handle._config.max_header_bytes, handle._config.max_header_count, handle._config.max_body_bytes, + handle._config.max_request_target_bytes, + preserve_trailing_data=True, ) primary_error: BaseException | None = None + route_error_event: RouteErrorEvent | None = None try: while not handle.closed: try: @@ -565,8 +708,33 @@ async def _connection_loop( return if request is None: continue + if _is_upgrade_attempt(request): + response = await self._dispatch_websocket( + task, + handle, + client, + request, + parser.trailing_data, + ) + if response is not None: + await self._send_response(task, handle, client, response) + return + if parser.trailing_data: + await self._send_response( + task, + handle, + client, + Response.text("pipelined requests are not supported", status=400), + ) + return try: response = await self.dispatch(request) + except RouteMatchTimeout as exc: + route_error_event = RouteErrorEvent( + route_id=exc.route_id, + category="route_match_timeout", + ) + response = Response.text("internal server error", status=500) except Exception: response = Response.text("internal server error", status=500) await self._send_response(task, handle, client, response) @@ -575,8 +743,45 @@ async def _connection_loop( primary_error = exc raise finally: + observer_channel = handle._route_observer_channel + if route_error_event is not None and observer_channel is not None: + observer_channel.enqueue(route_error_event, task) handle._connection_finished(task, client, primary_error) + async def _dispatch_websocket( + self, + task: Any, + handle: ServerHandle, + client: TransportHandle, + request: Request, + trailing_data: bytes, + ) -> Response | None: + route = self._websocket_routes.get(request.path) + if route is None: + return Response.text("not found", status=404) + invalid = _validate_upgrade(request, route) + if invalid is not None: + return invalid + if len(handle._websocket_states) >= min( + self._websocket_config.max_connections, + handle._config.max_connections, + ): + return Response.text("WebSocket capacity reached", status=503) + try: + state = _WebSocketState( + handle._runtime, + handle._transport, + client, + request, + route, + self._websocket_config, + trailing_data, + ) + except WebSocketUnavailable as exc: + return Response.text(str(exc), status=503) + await run_websocket_connection(task, state, handle) + return None + async def _send_response( self, task: Any, diff --git a/smallserver/http.py b/smallserver/http.py index 8c9dfa5..183dc96 100644 --- a/smallserver/http.py +++ b/smallserver/http.py @@ -10,7 +10,22 @@ from typing import Any _TOKEN = re.compile(r"^[!#$%&'*+.^_`|~0-9A-Za-z-]+$") -_REASONS = {200: "OK", 201: "Created", 204: "No Content", 400: "Bad Request", 404: "Not Found", 405: "Method Not Allowed", 413: "Payload Too Large", 500: "Internal Server Error", 503: "Service Unavailable"} +_REASONS = { + 101: "Switching Protocols", + 200: "OK", + 201: "Created", + 204: "No Content", + 400: "Bad Request", + 403: "Forbidden", + 404: "Not Found", + 405: "Method Not Allowed", + 413: "Payload Too Large", + 414: "URI Too Long", + 408: "Request Timeout", + 426: "Upgrade Required", + 500: "Internal Server Error", + 503: "Service Unavailable", +} class Headers(Mapping[str, str]): @@ -60,16 +75,45 @@ class Request: headers: Headers body: bytes = b"" version: str = "HTTP/1.1" + raw_target: str | None = None + query_string: str = "" + path_params: Mapping[str, str] = field(default_factory=dict) + route_pattern: str | None = None def __post_init__(self) -> None: if not _TOKEN.fullmatch(self.method): raise ValueError("invalid HTTP method") + if not isinstance(self.query_string, str): + raise TypeError("query_string must be a string") + if self.raw_target is None and "?" not in self.path and self.query_string: + raw_target = self.path + "?" + self.query_string + else: + raw_target = self.path if self.raw_target is None else self.raw_target + if not isinstance(raw_target, str) or not raw_target.startswith("/"): + raise ValueError("request path/target must start with '/'") + target_path, separator, target_query = raw_target.partition("?") + if self.raw_target is None: + if self.query_string and self.query_string != (target_query if separator else ""): + raise ValueError("request target fields are inconsistent") + object.__setattr__(self, "path", target_path) + object.__setattr__(self, "query_string", target_query if separator else "") + object.__setattr__(self, "raw_target", raw_target) + elif self.path != target_path or self.query_string != (target_query if separator else ""): + raise ValueError("request target fields are inconsistent") if not self.path.startswith("/"): raise ValueError("request path must start with '/'") if not isinstance(self.headers, Headers): object.__setattr__(self, "headers", Headers(self.headers)) if not isinstance(self.body, bytes): raise TypeError("request body must be bytes") + if not isinstance(self.path_params, Mapping): + raise TypeError("path_params must be a mapping") + params = dict(self.path_params) + if any(not isinstance(name, str) or not isinstance(value, str) for name, value in params.items()): + raise TypeError("path_params must map strings to strings") + object.__setattr__(self, "path_params", MappingProxyType(params)) + if self.route_pattern is not None and not isinstance(self.route_pattern, str): + raise TypeError("route_pattern must be a string or None") @dataclass(frozen=True) diff --git a/smallserver/routing.py b/smallserver/routing.py new file mode 100644 index 0000000..f8db1ed --- /dev/null +++ b/smallserver/routing.py @@ -0,0 +1,307 @@ +"""Deterministic static and timeout-bounded regular-expression routing.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Iterable, Mapping +from dataclasses import dataclass +import importlib +import math +import re +import time +from types import MappingProxyType +from typing import Any + +from .http import Request, Response + +Handler = Callable[[Request], Awaitable[Response]] +SUPPORTED_METHODS = frozenset({"GET", "POST", "PUT", "PATCH", "DELETE"}) + + +class RegexRoutesUnavailable(RuntimeError): + """Raised when regex routes are used without their optional dependency.""" + + +class RouteMatchTimeout(RuntimeError): + """Raised when a bounded regex route match exceeds its deadline.""" + + def __init__(self, route_id: str) -> None: + self.route_id = route_id + super().__init__("regular-expression route matching timed out ({})".format(route_id)) + + +class RoutePathTooLarge(RuntimeError): + """Raised before matching when a request path exceeds its routing bound.""" + + +@dataclass(frozen=True) +class RouteErrorEvent: + """Traceback-free, immutable routing failure data safe for observation.""" + + route_id: str + category: str + + +@dataclass(frozen=True) +class RegexRouteConfig: + """Finite limits applied to regex registration and hostile request paths.""" + + max_path_bytes: int = 8 * 1024 + max_pattern_length: int = 1024 + max_routes: int = 100 + match_timeout: float = 0.01 + total_match_timeout: float = 0.05 + max_named_captures: int = 20 + + def __post_init__(self) -> None: + for name in ("max_path_bytes", "max_pattern_length", "max_routes", "max_named_captures"): + value = getattr(self, name) + if type(value) is not int or value <= 0: + raise ValueError("{} must be a positive integer".format(name)) + for name in ("match_timeout", "total_match_timeout"): + value = getattr(self, name) + if type(value) not in (int, float) or not math.isfinite(value) or value <= 0: + raise ValueError("{} must be a finite positive number".format(name)) + + +@dataclass(frozen=True) +class RouteMatch: + handler: Handler | None + path_params: Mapping[str, str] + route_pattern: str | None + allowed_methods: tuple[str, ...] = () + + +@dataclass(frozen=True) +class _RegexRoute: + route_id: str + pattern: str + compiled: Any + handlers: Mapping[str, Handler] + + +class Router: + """Resolve static and ordered regex routes without slowing static lookup.""" + + def __init__(self, regex_config: RegexRouteConfig | None = None) -> None: + self._static: dict[tuple[str, str], Handler] = {} + self._regex: list[_RegexRoute] = [] + self._regex_by_pattern: dict[str, int] = {} + self.regex_config = regex_config or RegexRouteConfig() + + @staticmethod + def normalize_methods(methods: Iterable[str]) -> tuple[str, ...]: + try: + normalized = tuple(dict.fromkeys(method.upper() for method in methods)) + except AttributeError as exc: + raise ValueError("routes must use one or more supported HTTP methods") from exc + if not normalized or any(method not in SUPPORTED_METHODS for method in normalized): + raise ValueError("routes must use one or more supported HTTP methods") + return normalized + + def add_static(self, path: str, methods: tuple[str, ...], handler: Handler) -> None: + keys = [(method, path) for method in methods] + for key in keys: + if key in self._static: + raise ValueError("route already registered: {} {}".format(key[0], path)) + for key in keys: + self._static[key] = handler + + def static_handler(self, method: str, path: str) -> Handler | None: + """Return an exact static handler without allocating match context.""" + return self._static.get((method.upper(), path)) + + def add_regex(self, pattern: str, methods: tuple[str, ...], handler: Handler) -> None: + if not isinstance(pattern, str): + raise TypeError("regex route pattern must be a string") + existing_index = self._regex_by_pattern.get(pattern) + if existing_index is not None: + existing = self._regex[existing_index] + duplicate = next((method for method in methods if method in existing.handlers), None) + if duplicate is not None: + raise ValueError("regex route already registered: {}".format(duplicate)) + handlers = dict(existing.handlers) + handlers.update((method, handler) for method in methods) + self._regex[existing_index] = _RegexRoute( + existing.route_id, + existing.pattern, + existing.compiled, + MappingProxyType(handlers), + ) + return + + if len(self._regex) >= self.regex_config.max_routes: + raise ValueError("maximum registered regex routes exceeded") + compiled = self._compile(pattern) + route = _RegexRoute( + "regex-route-{}".format(len(self._regex) + 1), + pattern, + compiled, + MappingProxyType({method: handler for method in methods}), + ) + self._regex_by_pattern[pattern] = len(self._regex) + self._regex.append(route) + + def resolve(self, method: str, path: str) -> RouteMatch: + method = method.upper() + static = self._static.get((method, path)) + if static is not None: + return RouteMatch(static, MappingProxyType({}), path) + + if not self._regex: + allowed = tuple(sorted(method for method, registered_path in self._static if registered_path == path)) + return RouteMatch(None, MappingProxyType({}), None, allowed) + self._validate_path(path) + deadline = time.monotonic() + self.regex_config.total_match_timeout + cached: dict[int, Any] = {} + for index, route in enumerate(self._regex): + if method not in route.handlers: + continue + match = self._match(route, path, deadline) + cached[index] = match + if match is not None: + return RouteMatch( + route.handlers[method], + MappingProxyType(_captures(match)), + route.pattern, + ) + + allowed = {registered_method for registered_method, registered_path in self._static if registered_path == path} + for index, route in enumerate(self._regex): + match = cached.get(index) + if index not in cached: + match = self._match(route, path, deadline) + if match is not None: + allowed.update(route.handlers) + return RouteMatch(None, MappingProxyType({}), None, tuple(sorted(allowed))) + + def _compile(self, pattern: str) -> Any: + if not isinstance(pattern, str): + raise TypeError("regex route pattern must be a string") + if not pattern.startswith("/"): + raise ValueError("regex route pattern must start with a literal '/'") + if len(pattern) > self.regex_config.max_pattern_length: + raise ValueError("regex route pattern is too long") + try: + engine = importlib.import_module("regex") + except ImportError as exc: + raise RegexRoutesUnavailable( + "regular-expression routes require 'smallserver[regex-routes]'" + ) from exc + try: + compiled = engine.compile("(?=/)(?:{})".format(pattern)) + except Exception: + raise ValueError("invalid regex route pattern") from None + names = _named_group_names(pattern) + compiled_names = set(compiled.groupindex) + if set(names) != compiled_names: + raise ValueError("regex route named groups could not be validated") + if len(names) != len(compiled_names): + raise ValueError("regex route pattern contains duplicate named groups") + if len(compiled_names) > self.regex_config.max_named_captures: + raise ValueError("regex route pattern has too many named captures") + return compiled + + def _validate_path(self, path: str) -> None: + try: + size = len(path.encode("ascii")) + except UnicodeEncodeError as exc: + raise ValueError("request path must contain ASCII characters only") from exc + if size > self.regex_config.max_path_bytes: + raise RoutePathTooLarge("request path is too large for routing") + + def _match(self, route: _RegexRoute, path: str, deadline: float) -> Any: + remaining = deadline - time.monotonic() + if remaining <= 0: + raise RouteMatchTimeout(route.route_id) + timeout = min(float(self.regex_config.match_timeout), remaining) + try: + return route.compiled.fullmatch(path, timeout=timeout) + except TimeoutError as exc: + raise RouteMatchTimeout(route.route_id) from exc + + +def _captures(match: Any) -> dict[str, str]: + return {name: value for name, value in match.groupdict().items() if value is not None} + + +def _named_group_names(pattern: str) -> list[str]: + """Find declarations while honoring regex comments and scoped verbose mode.""" + names: list[str] = [] + verbose_stack = [False] + escaped = False + in_class = False + index = 0 + while index < len(pattern): + character = pattern[index] + if escaped: + escaped = False + index += 1 + continue + if character == "\\": + escaped = True + index += 1 + continue + if character == "[": + in_class = True + index += 1 + continue + if character == "]" and in_class: + in_class = False + index += 1 + continue + if not in_class and verbose_stack[-1] and character == "#": + newline = pattern.find("\n", index + 1) + index = len(pattern) if newline < 0 else newline + 1 + continue + if not in_class and pattern.startswith("(?#", index): + index = _comment_end(pattern, index + 3) + continue + if not in_class and character == "(": + flags = re.match(r"\(\?([A-Za-z]*)(?:-([A-Za-z]*))?([:)])", pattern[index:]) + if flags is not None: + enabled, disabled, delimiter = flags.groups() + verbose = (verbose_stack[-1] or "x" in enabled) and "x" not in (disabled or "") + index += flags.end() + if delimiter == ":": + verbose_stack.append(verbose) + else: + verbose_stack[-1] = verbose + continue + marker_length = 0 + if pattern.startswith("(?P<", index): + marker_length = 4 + elif pattern.startswith("(?<", index): + next_character = pattern[index + 3 : index + 4] + if next_character not in ("=", "!"): + marker_length = 3 + if marker_length: + end = pattern.find(">", index + marker_length) + if end >= 0: + names.append(pattern[index + marker_length : end]) + verbose_stack.append(verbose_stack[-1]) + index = end + 1 + continue + verbose_stack.append(verbose_stack[-1]) + index += 1 + continue + if not in_class and character == ")": + if len(verbose_stack) > 1: + verbose_stack.pop() + index += 1 + continue + index += 1 + return names + + +def _comment_end(pattern: str, index: int) -> int: + escaped = False + while index < len(pattern): + character = pattern[index] + if escaped: + escaped = False + elif character == "\\": + escaped = True + elif character == ")": + return index + 1 + index += 1 + return index diff --git a/smallserver/server.py b/smallserver/server.py index 36fdf21..cc87f74 100644 --- a/smallserver/server.py +++ b/smallserver/server.py @@ -2,13 +2,17 @@ from __future__ import annotations +from collections import deque from dataclasses import dataclass from typing import Any, Callable from ._transport import KernelTransport, TransportHandle, WakeupChannel from .http import Headers, Request, Response +from .routing import RouteErrorEvent from .runtime import ManagedRuntimeConfig +_ROUTE_OBSERVER_SIGNAL = 31 + class HTTPParseError(Exception): """A request rejected before it can be dispatched to application code.""" @@ -22,12 +26,27 @@ def __init__(self, status: int, detail: str) -> None: class HTTPRequestParser: """Incrementally parse one bounded HTTP/1.1 request with Content-Length.""" - def __init__(self, max_header_bytes: int, max_header_count: int, max_body_bytes: int) -> None: + def __init__( + self, + max_header_bytes: int, + max_header_count: int, + max_body_bytes: int, + max_request_target_bytes: int = 8 * 1024, + preserve_trailing_data: bool = False, + ) -> None: self._max_header_bytes = max_header_bytes self._max_header_count = max_header_count self._max_body_bytes = max_body_bytes + self._max_request_target_bytes = max_request_target_bytes + self._preserve_trailing_data = preserve_trailing_data self._buffer = bytearray() self._request_head: tuple[str, str, Headers, int] | None = None + self._trailing_data = b"" + + @property + def trailing_data(self) -> bytes: + """Bytes received after the request body for an explicit protocol handoff.""" + return self._trailing_data def feed(self, data: bytes) -> Request | None: self._buffer.extend(data) @@ -43,13 +62,24 @@ def feed(self, data: bytes) -> Request | None: self._request_head = self._parse_head(bytes(self._buffer[:marker])) del self._buffer[:header_length] - method, path, headers, content_length = self._request_head - if len(self._buffer) > content_length: + method, raw_target, headers, content_length = self._request_head + if len(self._buffer) > content_length and not self._preserve_trailing_data: raise HTTPParseError(400, "pipelined requests are not supported") if len(self._buffer) < content_length: return None try: - return Request(method, path, headers, bytes(self._buffer), "HTTP/1.1") + path, separator, query_string = raw_target.partition("?") + body = bytes(self._buffer[:content_length]) + self._trailing_data = bytes(self._buffer[content_length:]) + return Request( + method, + path, + headers, + body, + "HTTP/1.1", + raw_target=raw_target, + query_string=query_string if separator else "", + ) except ValueError as exc: raise HTTPParseError(400, str(exc)) from exc @@ -60,10 +90,12 @@ def _parse_head(self, raw: bytes) -> tuple[str, str, Headers, int]: raise HTTPParseError(400, "request headers are not valid bytes") from exc if not lines or len(lines[0].split(" ")) != 3: raise HTTPParseError(400, "malformed request line") - method, path, version = lines[0].split(" ") - if version != "HTTP/1.1" or not path.startswith("/"): + method, raw_target, version = lines[0].split(" ") + if version != "HTTP/1.1" or not raw_target.startswith("/"): raise HTTPParseError(400, "only origin-form HTTP/1.1 requests are supported") - if "#" in path or any(not 0x21 <= ord(character) <= 0x7E for character in path): + if len(raw_target.encode("iso-8859-1")) > self._max_request_target_bytes: + raise HTTPParseError(414, "request target is too large") + if "#" in raw_target or any(not 0x21 <= ord(character) <= 0x7E for character in raw_target): raise HTTPParseError(400, "request target is not valid origin-form") if len(lines) - 1 > self._max_header_count: raise HTTPParseError(413, "too many request headers") @@ -95,7 +127,7 @@ def _parse_head(self, raw: bytes) -> tuple[str, str, Headers, int]: raise HTTPParseError(413, "request body is too large") if not headers.get("host"): raise HTTPParseError(400, "HTTP/1.1 requests require a Host header") - return method, path, headers, length + return method, raw_target, headers, length @dataclass(frozen=True) @@ -110,6 +142,8 @@ class ServerConfig: listener_priority: int = 1 connection_priority: int = 2 accept_batch_size: int = 16 + max_request_target_bytes: int = 8 * 1024 + max_route_error_events: int = 16 managed_runtime: ManagedRuntimeConfig | None = None def __post_init__(self) -> None: @@ -122,6 +156,8 @@ def __post_init__(self) -> None: "listener_priority", "connection_priority", "accept_batch_size", + "max_request_target_bytes", + "max_route_error_events", ): value = getattr(self, name) if type(value) is not int or value <= 0: @@ -132,6 +168,67 @@ def __post_init__(self) -> None: raise TypeError("managed_runtime must be a ManagedRuntimeConfig or None") +class RouteObserverChannel: + """Bounded scheduler-local delivery state for one server invocation.""" + + def __init__(self, observer: Any, max_events: int) -> None: + self.observer = observer + self.max_events = max_events + self.events: deque[RouteErrorEvent] = deque() + self.task: Any = None + self.accepting = True + self.dropped = 0 + self.failures = 0 + + def bind(self, task: Any) -> None: + self.task = task + + def enqueue(self, event: RouteErrorEvent, source_task: Any) -> bool: + if not self.accepting or len(self.events) >= self.max_events: + self.dropped += 1 + return False + self.events.append(event) + try: + signalled = ( + self.task is not None + and source_task.sendSignal(self.task.getID(), _ROUTE_OBSERVER_SIGNAL) == 0 + ) + except BaseException: + signalled = False + if not signalled: + self.events.pop() + self.dropped += 1 + return False + return True + + def stop(self) -> None: + self.accepting = False + self.dropped += len(self.events) + self.events.clear() + task = self.task + if task is not None and not getattr(task, "done", False): + try: + task.acceptSignal(_ROUTE_OBSERVER_SIGNAL) + except BaseException: + pass + + +async def run_route_observer(task: Any, channel: RouteObserverChannel) -> None: + """Drain sanitized events on a dedicated SmallOS task.""" + try: + while channel.accepting: + while channel.events: + event = channel.events.popleft() + try: + channel.observer(event) + except BaseException: + channel.failures += 1 + if channel.accepting: + await task.wait_signal(_ROUTE_OBSERVER_SIGNAL) + finally: + channel.task = None + + class ServerHandle: """A bound listener and its cooperative shutdown signal.""" @@ -145,12 +242,14 @@ def __init__( wakeup: WakeupChannel | None, config: ServerConfig, on_finalized: Callable[[ServerHandle], None] | None = None, + route_observer_channel: RouteObserverChannel | None = None, ) -> None: self._runtime = runtime self._transport = transport self._listener = listener self._wakeup = wakeup self._config = config + self._route_observer_channel = route_observer_channel self._address = transport.local_address(listener) self._on_finalized = on_finalized self._close_requested = False @@ -167,6 +266,7 @@ def __init__( self._connections: dict[int, tuple[TransportHandle, Any]] = {} self._closing_connections: dict[int, TransportHandle] = {} self._pending_task_cancellations: dict[int, Any] = {} + self._websocket_states: dict[int, Any] = {} self._capacity_waiting = False @property @@ -225,6 +325,16 @@ def _notify_capacity_released(self, previous_count: int) -> None: except BaseException as error: self._listener_failed(error, getattr(self._runtime, "cursor", None)) + @property + def dropped_route_error_events(self) -> int: + channel = self._route_observer_channel + return 0 if channel is None else int(channel.dropped) + + @property + def route_observer_failures(self) -> int: + channel = self._route_observer_channel + return 0 if channel is None else int(channel.failures) + def close(self) -> None: """Request external shutdown through a kernel wakeup channel.""" if self._finished: @@ -263,6 +373,10 @@ async def close_from_task(self, task: Any) -> None: return self._close_requested = True self._finalization_attempted = True + for state in tuple(self._websocket_states.values()): + request_shutdown = getattr(state, "request_shutdown", None) + if callable(request_shutdown): + request_shutdown() self._finish_close(current_task=task) def _listener_failed(self, exc: BaseException, task: Any) -> None: @@ -297,6 +411,13 @@ def _finish_close( return self._close_requested = True self._finalization_attempted = True + for state in tuple(self._websocket_states.values()): + request_shutdown = getattr(state, "request_shutdown", None) + if callable(request_shutdown): + request_shutdown() + channel = self._route_observer_channel + if channel is not None: + channel.stop() if self._wakeup is not None and not self._wakeup.closed: try: self._wakeup.close() @@ -327,6 +448,7 @@ def _finish_close( self._cancelled_task_ids.add(identity) for identity, (connection, task) in list(self._connections.items()): + websocket_owned = identity in self._websocket_states if owner_thread: if ( task is not current_task @@ -341,6 +463,10 @@ def _finish_close( self._runtime.resume_task(task) except BaseException: pass + if websocket_owned: + # The live coordinator owns its bounded Close handshake + # and releases the stream through _connection_finished(). + continue if task is not current_task: self._connections.pop(identity, None) self._close_or_retain(connection, current_task) @@ -471,6 +597,7 @@ def _update_finished(self) -> None: and not self._connections and not self._closing_connections and not self._pending_task_cancellations + and not self._websocket_states ) if self._finished: self._cleanup_errors.clear() diff --git a/smallserver/websocket.py b/smallserver/websocket.py new file mode 100644 index 0000000..9e024e3 --- /dev/null +++ b/smallserver/websocket.py @@ -0,0 +1,1077 @@ +"""Bounded RFC 6455 WebSocket support driven by SmallOS tasks.""" + +from __future__ import annotations + +import base64 +from collections import deque +from dataclasses import dataclass +import hashlib +import importlib +import inspect +import math +import time +from typing import Any, AsyncIterator, Callable, Mapping + +from .http import Headers, Request, Response + +_GUID = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11" +_DECISION_SIGNAL = 24 +_INBOX_SIGNAL = 25 +_OUTBOX_SIGNAL = 26 +_ACK_SIGNAL = 27 +_HTTP_TOKEN_CHARACTERS = frozenset( + "!#$%&'*+-.^_`|~0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +) + + +class WebSocketUnavailable(RuntimeError): + """The optional WebSocket protocol dependency is unavailable.""" + + +class WebSocketStateError(RuntimeError): + """A WebSocket operation is invalid in the current lifecycle state.""" + + +class WebSocketCapacityError(RuntimeError): + """A bounded WebSocket mailbox cannot accept another item.""" + + +class WebSocketDisconnect(Exception): + """The peer or server closed a WebSocket connection.""" + + def __init__(self, code: int = 1006, reason: str = "") -> None: + self.code = code + self.reason = reason + super().__init__("WebSocket disconnected ({})".format(code)) + + +@dataclass(frozen=True) +class WebSocketMessage: + """One complete text or binary WebSocket message.""" + + data: str | bytes + + @property + def is_text(self) -> bool: + return isinstance(self.data, str) + + @property + def is_binary(self) -> bool: + return isinstance(self.data, bytes) + + @property + def text(self) -> str: + if not isinstance(self.data, str): + raise TypeError("WebSocket message is binary") + return self.data + + @property + def bytes(self) -> bytes: + if not isinstance(self.data, bytes): + raise TypeError("WebSocket message is text") + return self.data + + +@dataclass(frozen=True) +class WebSocketConfig: + """Finite resource and lifetime limits for WebSocket connections.""" + + max_frame_payload_bytes: int = 1024 * 1024 + max_message_bytes: int = 1024 * 1024 + max_inbound_messages: int = 16 + max_inbound_bytes: int = 2 * 1024 * 1024 + max_outbound_commands: int = 16 + max_outbound_bytes: int = 2 * 1024 * 1024 + receive_chunk_bytes: int = 16 * 1024 + write_chunk_bytes: int = 16 * 1024 + max_connections: int = 100 + handshake_timeout: float = 10.0 + idle_timeout: float = 300.0 + pong_timeout: float = 10.0 + write_timeout: float = 30.0 + close_timeout: float = 5.0 + deadline_resolution: float = 0.05 + + def __post_init__(self) -> None: + integer_fields = ( + "max_frame_payload_bytes", + "max_message_bytes", + "max_inbound_messages", + "max_inbound_bytes", + "max_outbound_commands", + "max_outbound_bytes", + "receive_chunk_bytes", + "write_chunk_bytes", + "max_connections", + ) + for name in integer_fields: + value = getattr(self, name) + if type(value) is not int or value <= 0: + raise ValueError("{} must be a positive integer".format(name)) + for name in ( + "handshake_timeout", + "idle_timeout", + "pong_timeout", + "write_timeout", + "close_timeout", + "deadline_resolution", + ): + value = getattr(self, name) + if type(value) not in (int, float) or not math.isfinite(value) or value <= 0: + raise ValueError("{} must be a finite positive number".format(name)) + if self.max_frame_payload_bytes > self.max_message_bytes: + raise ValueError("max_frame_payload_bytes must not exceed max_message_bytes") + + +@dataclass(frozen=True) +class _WebSocketRoute: + handler: Callable[[WebSocket], Any] + origins: frozenset[str] | None + subprotocols: tuple[str, ...] + + +@dataclass +class _OutboundCommand: + event: Any + size: int + waiter: Any = None + done: bool = False + error: BaseException | None = None + + +@dataclass(frozen=True) +class _WSProtoAPI: + Connection: Any + ConnectionType: Any + TextMessage: Any + BytesMessage: Any + Ping: Any + Pong: Any + CloseConnection: Any + + +def _load_wsproto() -> _WSProtoAPI: + """Import the optional protocol engine only when an upgrade is served.""" + try: + connection = importlib.import_module("wsproto.connection") + events = importlib.import_module("wsproto.events") + except ImportError as exc: + raise WebSocketUnavailable( + "WebSocket routes require the 'smallserver[websocket]' extra" + ) from exc + return _WSProtoAPI( + connection.Connection, + connection.ConnectionType, + events.TextMessage, + events.BytesMessage, + events.Ping, + events.Pong, + events.CloseConnection, + ) + + +class _FrameGuard: + """Validate declared client frame sizes before forwarding payload bytes.""" + + def __init__(self, max_payload_bytes: int) -> None: + self._max_payload_bytes = max_payload_bytes + self._header = bytearray() + self._header_length: int | None = None + self._payload_remaining = 0 + + def feed(self, data: bytes) -> tuple[bytes, ...]: + chunks: list[bytes] = [] + offset = 0 + while offset < len(data): + if self._payload_remaining: + length = min(self._payload_remaining, len(data) - offset) + chunks.append(data[offset : offset + length]) + offset += length + self._payload_remaining -= length + if self._payload_remaining == 0: + self._header_length = None + continue + + if self._header_length is None: + needed = 2 - len(self._header) + if needed: + length = min(needed, len(data) - offset) + self._header.extend(data[offset : offset + length]) + offset += length + if len(self._header) < 2: + continue + second = self._header[1] + marker = second & 0x7F + extension = 2 if marker == 126 else 8 if marker == 127 else 0 + self._header_length = 2 + extension + (4 if second & 0x80 else 0) + + needed = self._header_length - len(self._header) + if needed: + length = min(needed, len(data) - offset) + self._header.extend(data[offset : offset + length]) + offset += length + if len(self._header) < self._header_length: + continue + + first, second = self._header[:2] + if not second & 0x80: + raise ValueError("client WebSocket frames must be masked") + marker = second & 0x7F + index = 2 + if marker == 126: + payload_length = int.from_bytes(self._header[index : index + 2], "big") + index += 2 + if payload_length < 126: + raise ValueError("non-minimal WebSocket frame length") + elif marker == 127: + payload_length = int.from_bytes(self._header[index : index + 8], "big") + index += 8 + if payload_length < 65536 or payload_length >> 63: + raise ValueError("invalid WebSocket frame length") + else: + payload_length = marker + opcode = first & 0x0F + if opcode >= 8 and (not first & 0x80 or payload_length > 125): + raise ValueError("invalid WebSocket control frame") + if payload_length > self._max_payload_bytes: + raise WebSocketCapacityError("WebSocket frame payload is too large") + chunks.append(bytes(self._header)) + self._header.clear() + self._payload_remaining = payload_length + if payload_length == 0: + self._header_length = None + return tuple(chunks) + + +class WebSocket: + """Application-facing WebSocket connection.""" + + def __init__(self, state: _WebSocketState) -> None: + self._state = state + + @property + def request(self) -> Request: + return self._state.request + + @property + def subprotocol(self) -> str | None: + return self._state.subprotocol + + async def accept( + self, + subprotocol: str | None = None, + headers: Mapping[str, str] | None = None, + ) -> None: + await self._state.accept(subprotocol, headers) + + async def reject(self, response: Response) -> None: + await self._state.reject(response) + + async def receive(self) -> WebSocketMessage: + return await self._state.receive() + + async def receive_text(self) -> str: + return (await self.receive()).text + + async def receive_bytes(self) -> bytes: + return (await self.receive()).bytes + + async def send_text(self, value: str) -> None: + if not isinstance(value, str): + raise TypeError("WebSocket text payload must be a string") + await self._state.send_message(value) + + async def send_bytes(self, value: bytes) -> None: + if not isinstance(value, bytes): + raise TypeError("WebSocket binary payload must be bytes") + await self._state.send_message(value) + + async def ping(self, payload: bytes = b"") -> None: + if not isinstance(payload, bytes) or len(payload) > 125: + raise ValueError("WebSocket Ping payload must be at most 125 bytes") + await self._state.ping(payload) + + async def close(self, code: int = 1000, reason: str = "") -> None: + await self._state.close(code, reason) + + def __aiter__(self) -> AsyncIterator[WebSocketMessage]: + return self + + async def __anext__(self) -> WebSocketMessage: + try: + return await self.receive() + except WebSocketDisconnect as exc: + raise StopAsyncIteration from exc + + +class _WebSocketState: + """One coordinator-owned WebSocket protocol and mailbox state.""" + + def __init__( + self, + runtime: Any, + transport: Any, + client: Any, + request: Request, + route: _WebSocketRoute, + config: WebSocketConfig, + trailing_data: bytes, + ) -> None: + self.runtime = runtime + self.transport = transport + self.client = client + self.request = request + self.route = route + self.config = config + self.trailing_data = trailing_data + self.api = _load_wsproto() + self.protocol: Any = None + self.coordinator_task: Any = None + self.handler_task: Any = None + self.reader_task: Any = None + self.writer_task: Any = None + self.deadline_task: Any = None + self.accepted = False + self.rejected = False + self.handshake_state = "pending" + self.handshake_error: BaseException | None = None + self.shutdown = False + self.peer_closed = False + self.close_sent = False + self.subprotocol: str | None = None + self.disconnect: WebSocketDisconnect | None = None + self.handler_error: BaseException | None = None + self.fatal_error: BaseException | None = None + self.inbox: deque[WebSocketMessage] = deque() + self.inbox_bytes = 0 + self.outbox: deque[_OutboundCommand] = deque() + self.outbox_bytes = 0 + self.writer_busy = False + self.active_command: _OutboundCommand | None = None + self._message_kind: type | None = None + self._message_buffer = bytearray() + self._message_bytes = 0 + self._guard = _FrameGuard(config.max_frame_payload_bytes) + self.created_at = time.monotonic() + self.last_activity = self.created_at + self.pong_deadline: float | None = None + self._ping_generation = 0 + self._pending_ping_generation: int | None = None + self._pending_ping_payload: bytes | None = None + self.close_deadline: float | None = None + self.write_deadline: float | None = None + self._children: list[Any] = [] + + def _current_task(self) -> Any: + task = getattr(self.runtime, "cursor", None) + if task is None: + raise WebSocketStateError("WebSocket operations require a running SmallOS task") + return task + + @staticmethod + def _signal(task: Any, signal: int) -> None: + if task is not None: + accept = getattr(task, "acceptSignal", None) + if callable(accept): + accept(signal) + + async def _send_http(self, task: Any, response: Response) -> None: + headers = { + name: value + for name, value in response.headers.items() + if name.lower() != "connection" + } + headers["Connection"] = "close" + payload = Response(response.status, response.body, headers).to_http1() + await self.transport.send_all(task, self.client, payload) + + async def accept( + self, subprotocol: str | None, headers: Mapping[str, str] | None + ) -> None: + task = self._current_task() + if self.handshake_state != "pending": + raise WebSocketStateError("WebSocket handshake is already decided") + if subprotocol is not None: + if subprotocol not in self.route.subprotocols: + raise WebSocketStateError("selected subprotocol is not allowed by the route") + if subprotocol not in _token_list(self.request.headers.get("sec-websocket-protocol")): + raise WebSocketStateError("selected subprotocol was not offered by the client") + extra = Headers(headers or {}) + forbidden = { + "connection", + "upgrade", + "sec-websocket-accept", + "sec-websocket-protocol", + "sec-websocket-extensions", + "content-length", + } + if any(name.lower() in forbidden for name in extra): + raise ValueError("handshake headers contain a reserved field") + key = self.request.headers["sec-websocket-key"].encode("ascii") + accept_value = base64.b64encode(hashlib.sha1(key + _GUID).digest()).decode("ascii") + lines = [ + "HTTP/1.1 101 Switching Protocols", + "Upgrade: websocket", + "Connection: Upgrade", + "Sec-WebSocket-Accept: {}".format(accept_value), + ] + if subprotocol is not None: + lines.append("Sec-WebSocket-Protocol: {}".format(subprotocol)) + lines.extend("{}: {}".format(name, value) for name, value in extra.items()) + payload = ("\r\n".join(lines) + "\r\n\r\n").encode("latin-1") + protocol = self.api.Connection(self.api.ConnectionType.SERVER) + self.handshake_state = "accepting" + self.protocol = protocol + try: + await self.transport.send_all(task, self.client, payload) + except GeneratorExit: + raise + except BaseException as exc: + self._fail_handshake(exc) + raise + else: + self.subprotocol = subprotocol + self.accepted = True + self.handshake_state = "accepted" + self.last_activity = time.monotonic() + self._signal(self.coordinator_task, _DECISION_SIGNAL) + + async def reject(self, response: Response) -> None: + task = self._current_task() + if self.handshake_state != "pending": + raise WebSocketStateError("WebSocket handshake is already decided") + if not isinstance(response, Response): + raise TypeError("reject() requires a Response") + if response.status < 300: + raise ValueError("WebSocket rejection response must have status 300 or greater") + self.handshake_state = "rejecting" + try: + await self._send_http(task, response) + except GeneratorExit: + raise + except BaseException as exc: + self._fail_handshake(exc) + raise + else: + self.rejected = True + self.handshake_state = "rejected" + self._signal(self.coordinator_task, _DECISION_SIGNAL) + + def _fail_handshake(self, error: BaseException) -> None: + if self.handshake_state == "failed": + return + self.accepted = False + self.rejected = False + self.protocol = None + self.handshake_state = "failed" + self.handshake_error = error + if isinstance(error, (KeyboardInterrupt, SystemExit)): + self.fatal_error = error + self._disconnect(1006) + + def _require_open(self) -> None: + if not self.accepted: + raise WebSocketStateError("WebSocket must be accepted first") + if self.shutdown or self.disconnect is not None: + raise self.disconnect or WebSocketStateError("WebSocket is closed") + + async def receive(self) -> WebSocketMessage: + task = self._current_task() + if not self.accepted: + raise WebSocketStateError("WebSocket must be accepted first") + while not self.inbox: + if self.disconnect is not None: + raise self.disconnect + await task.wait_signal(_INBOX_SIGNAL) + message = self.inbox.popleft() + self.inbox_bytes -= _message_size(message.data) + return message + + async def send_message(self, value: str | bytes) -> None: + self._require_open() + size = _message_size(value) + if size > self.config.max_message_bytes: + raise WebSocketCapacityError("WebSocket message is too large") + event = ( + self.api.TextMessage(data=value) + if isinstance(value, str) + else self.api.BytesMessage(data=value) + ) + await self._enqueue(event, size, wait=True) + + async def ping(self, payload: bytes) -> None: + self._require_open() + if self._pending_ping_generation is not None: + raise WebSocketStateError("a WebSocket Ping is already awaiting Pong") + self._ping_generation += 1 + generation = self._ping_generation + self._pending_ping_generation = generation + self._pending_ping_payload = payload + self.pong_deadline = time.monotonic() + self.config.pong_timeout + try: + await self._enqueue( + self.api.Ping(payload=payload), len(payload), wait=True + ) + except BaseException: + if self._pending_ping_generation == generation: + self._clear_pending_ping() + raise + + async def close(self, code: int, reason: str) -> None: + self._require_open() + if not _valid_close_code(code): + raise ValueError("invalid WebSocket close code") + if not isinstance(reason, str) or len(reason.encode("utf-8")) > 123: + raise ValueError("WebSocket close reason must be at most 123 UTF-8 bytes") + reason_bytes = reason.encode("utf-8") + self.close_sent = True + self.close_deadline = time.monotonic() + self.config.close_timeout + try: + await self._enqueue( + self.api.CloseConnection(code=code, reason=reason), + 2 + len(reason_bytes), + wait=True, + ) + except BaseException: + self._disconnect(code, reason) + raise + task = self._current_task() + while not self.peer_closed and time.monotonic() < self.close_deadline: + await task.sleep(min(self.config.deadline_resolution, self.config.close_timeout)) + if self.disconnect is None: + self.disconnect = WebSocketDisconnect(code, reason) + + async def _enqueue(self, event: Any, size: int, *, wait: bool) -> None: + if ( + len(self.outbox) >= self.config.max_outbound_commands + or self.outbox_bytes + size > self.config.max_outbound_bytes + ): + raise WebSocketCapacityError("WebSocket outbound queue is full") + waiter = self._current_task() if wait else None + command = _OutboundCommand(event, size, waiter) + self.outbox.append(command) + self.outbox_bytes += size + self._signal(self.writer_task, _OUTBOX_SIGNAL) + if wait: + while not command.done: + await waiter.wait_signal(_ACK_SIGNAL) + if command.error is not None: + raise command.error + + def _clear_pending_ping(self) -> None: + self._pending_ping_generation = None + self._pending_ping_payload = None + self.pong_deadline = None + + def _fail_outbound(self, error: BaseException) -> None: + commands = list(self.outbox) + self.outbox.clear() + self.outbox_bytes = 0 + if self.active_command is not None: + commands.insert(0, self.active_command) + seen: set[int] = set() + for command in commands: + if id(command) in seen: + continue + seen.add(id(command)) + command.error = error + command.done = True + self._signal(command.waiter, _ACK_SIGNAL) + + def _cancel_task(self, target: Any) -> None: + if target is None or target is getattr(self.runtime, "cursor", None): + return + try: + self.runtime.cancel_task(target) + except (KeyboardInterrupt, SystemExit) as exc: + self.fatal_error = exc + raise + except BaseException: + # ServerHandle retains ownership and retries cancellation in finalization. + pass + + def _abort_writer(self, error: BaseException) -> None: + self._fail_outbound(error) + self._cancel_task(self.writer_task) + + def _cancel_handler(self) -> None: + self._cancel_task(self.handler_task) + + def _enqueue_control(self, event: Any, size: int = 0) -> bool: + if ( + len(self.outbox) >= self.config.max_outbound_commands + or self.outbox_bytes + size > self.config.max_outbound_bytes + ): + return False + self.outbox.append(_OutboundCommand(event, size)) + self.outbox_bytes += size + self._signal(self.writer_task, _OUTBOX_SIGNAL) + return True + + def _deliver_message(self, value: str | bytes) -> bool: + size = _message_size(value) + if ( + len(self.inbox) >= self.config.max_inbound_messages + or self.inbox_bytes + size > self.config.max_inbound_bytes + ): + return False + self.inbox.append(WebSocketMessage(value)) + self.inbox_bytes += size + self._signal(self.handler_task, _INBOX_SIGNAL) + return True + + def _disconnect(self, code: int, reason: str = "") -> None: + if self.disconnect is None: + self.disconnect = WebSocketDisconnect(code, reason) + self._signal(self.handler_task, _INBOX_SIGNAL) + self._signal(self.coordinator_task, _DECISION_SIGNAL) + + def request_shutdown(self, code: int = 1001) -> None: + if self.shutdown: + return + self.shutdown = True + if self.accepted and not self.close_sent and self.protocol is not None: + self._enqueue_control( + self.api.CloseConnection(code=code, reason="server shutdown"), 17 + ) + self.close_sent = True + if self.handshake_state in {"pending", "accepting", "rejecting"}: + self._fail_handshake(WebSocketDisconnect(code, "server shutdown")) + else: + self._disconnect(code, "server shutdown") + self._cancel_handler() + self._signal(self.writer_task, _OUTBOX_SIGNAL) + + +def _token_list(value: str | None) -> tuple[str, ...]: + if value is None: + return () + return tuple(token.strip() for token in value.split(",") if token.strip()) + + +def _is_http_token(value: str) -> bool: + return ( + isinstance(value, str) + and bool(value) + and all(character in _HTTP_TOKEN_CHARACTERS for character in value) + ) + + +def _message_size(value: str | bytes) -> int: + return len(value.encode("utf-8")) if isinstance(value, str) else len(value) + + +def _valid_close_code(code: int) -> bool: + return type(code) is int and ( + 1000 <= code <= 1014 and code not in {1004, 1005, 1006} + or 3000 <= code <= 4999 + ) + + +def _valid_websocket_key(value: str | None) -> bool: + if value is None: + return False + try: + return len(base64.b64decode(value.encode("ascii"), validate=True)) == 16 + except (ValueError, UnicodeEncodeError): + return False + + +def _is_upgrade_attempt(request: Request) -> bool: + headers = request.headers + return bool( + headers.get("upgrade") + or headers.get("sec-websocket-key") + or headers.get("sec-websocket-version") + or any(token.lower() == "upgrade" for token in _token_list(headers.get("connection"))) + ) + + +def _validate_upgrade( + request: Request, route: _WebSocketRoute +) -> Response | None: + if request.method.upper() != "GET": + return Response.text("WebSocket upgrade requires GET", status=400) + if request.version != "HTTP/1.1": + return Response.text("WebSocket upgrade requires HTTP/1.1", status=400) + if request.body: + return Response.text("WebSocket upgrade must not include a body", status=400) + if request.headers.get("upgrade", "").lower() != "websocket": + return Response.text("invalid WebSocket Upgrade header", status=400) + connection_tokens = { + token.lower() for token in _token_list(request.headers.get("connection")) + } + if "upgrade" not in connection_tokens: + return Response.text("invalid WebSocket Connection header", status=400) + if request.headers.get("sec-websocket-version") != "13": + return Response.text( + "unsupported WebSocket version", + status=426, + headers={"Sec-WebSocket-Version": "13"}, + ) + if not _valid_websocket_key(request.headers.get("sec-websocket-key")): + return Response.text("invalid WebSocket key", status=400) + offered_subprotocols = _token_list( + request.headers.get("sec-websocket-protocol") + ) + offered_header = request.headers.get("sec-websocket-protocol") + if offered_header is not None and ( + not offered_subprotocols + or any(not _is_http_token(protocol) for protocol in offered_subprotocols) + ): + return Response.text("invalid WebSocket subprotocol", status=400) + origin = request.headers.get("origin") + if route.origins is not None and origin not in route.origins: + return Response.text("WebSocket origin is not allowed", status=403) + return None + + +async def run_websocket_connection( + task: Any, state: _WebSocketState, server_handle: Any +) -> None: + """Coordinate one accepted HTTP connection through its WebSocket lifetime.""" + state.coordinator_task = task + server_handle._websocket_states[id(state.client)] = state + + def spawn(routine: Any, name: str) -> Any: + child = task.spawn( + routine, + priority=server_handle._config.connection_priority, + args=(state,), + name=name, + ) + state._children.append(child) + server_handle._owned_tasks.append(child) + return child + + try: + state.handler_task = spawn(_run_handler, "smallserver-websocket-handler") + state.deadline_task = spawn( + _run_deadlines, "smallserver-websocket-deadline" + ) + while state.handshake_state in {"pending", "accepting", "rejecting"}: + await task.wait_signal(_DECISION_SIGNAL) + if not state.accepted: + if state.fatal_error is not None: + raise state.fatal_error + return + state.writer_task = spawn(_run_writer, "smallserver-websocket-writer") + state.reader_task = spawn(_run_reader, "smallserver-websocket-reader") + try: + await task.join(state.handler_task) + except WebSocketDisconnect: + pass + except BaseException as exc: + state.handler_error = exc + + if isinstance(state.handler_error, (KeyboardInterrupt, SystemExit)): + raise state.handler_error + if state.fatal_error is not None: + raise state.fatal_error + + if state.handler_error is not None and not state.close_sent: + try: + state.close_deadline = time.monotonic() + state.config.close_timeout + await state._enqueue( + state.api.CloseConnection(code=1011, reason="handler failed"), + 16, + wait=True, + ) + state.close_sent = True + except (KeyboardInterrupt, SystemExit): + raise + except BaseException: + pass + elif not state.close_sent and state.disconnect is None: + try: + state.close_deadline = time.monotonic() + state.config.close_timeout + await state._enqueue( + state.api.CloseConnection(code=1000, reason=""), 2, wait=True + ) + state.close_sent = True + except (KeyboardInterrupt, SystemExit): + raise + except BaseException: + pass + + deadline = time.monotonic() + state.config.close_timeout + while ( + state.outbox or state.writer_busy or not state.peer_closed + ) and time.monotonic() < deadline: + await task.sleep( + min(state.config.deadline_resolution, state.config.close_timeout) + ) + finally: + state.request_shutdown() + for child in list(state._children): + if child is task: + continue + if ( + server_handle._cancel_or_retain_task(child) + and child in server_handle._owned_tasks + ): + server_handle._owned_tasks.remove(child) + state._children.clear() + server_handle._websocket_states.pop(id(state.client), None) + + +async def _run_handler(task: Any, state: _WebSocketState) -> None: + socket = WebSocket(state) + try: + result = state.route.handler(socket) + if not inspect.isawaitable(result): + raise TypeError("WebSocket handlers must return an awaitable") + await result + except WebSocketDisconnect: + pass + except GeneratorExit: + raise + except (KeyboardInterrupt, SystemExit) as exc: + state.handler_error = exc + if state.handshake_state in {"pending", "accepting", "rejecting"}: + state._fail_handshake(exc) + raise + except BaseException as exc: + state.handler_error = exc + finally: + if state.handshake_state == "pending": + response = Response.text( + "internal server error" if state.handler_error is not None else "forbidden", + status=500 if state.handler_error is not None else 403, + ) + try: + await state.reject(response) + except (KeyboardInterrupt, SystemExit) as exc: + state.handler_error = exc + state.fatal_error = exc + raise + except BaseException: + state._disconnect(1006) + state._signal(state.coordinator_task, _DECISION_SIGNAL) + + +async def _run_writer(task: Any, state: _WebSocketState) -> None: + while not state.shutdown or state.outbox: + while state.outbox: + command = state.outbox.popleft() + state.outbox_bytes -= command.size + state.writer_busy = True + state.active_command = command + state.write_deadline = time.monotonic() + state.config.write_timeout + try: + payload = state.protocol.send(command.event) + for offset in range(0, len(payload), state.config.write_chunk_bytes): + await state.transport.send_all( + task, + state.client, + payload[offset : offset + state.config.write_chunk_bytes], + ) + state.last_activity = time.monotonic() + except GeneratorExit: + raise + except (KeyboardInterrupt, SystemExit) as exc: + command.error = exc + state.fatal_error = exc + state._disconnect(1006) + state._cancel_handler() + raise + except BaseException as exc: + command.error = exc + state._disconnect(1006) + state._cancel_handler() + finally: + state.writer_busy = False + if state.active_command is command: + state.active_command = None + state.write_deadline = None + command.done = True + state._signal(command.waiter, _ACK_SIGNAL) + if not state.shutdown: + await task.wait_signal(_OUTBOX_SIGNAL) + + +async def _run_reader(task: Any, state: _WebSocketState) -> None: + try: + if state.trailing_data: + _receive_protocol_data(state, state.trailing_data) + state.trailing_data = b"" + if _drain_protocol_events(state): + return + while not state.shutdown: + chunk = await state.transport.recv( + task, state.client, state.config.receive_chunk_bytes + ) + if not chunk: + state.protocol.receive_data(None) + _drain_protocol_events(state) + state._disconnect(1006) + state._cancel_handler() + return + state.last_activity = time.monotonic() + _receive_protocol_data(state, chunk) + if _drain_protocol_events(state): + return + except GeneratorExit: + raise + except WebSocketCapacityError: + if state.shutdown: + return + state._enqueue_control( + state.api.CloseConnection(code=1009, reason="message too large"), 19 + ) + state.close_sent = True + state._disconnect(1009, "message too large") + state._cancel_handler() + except (KeyboardInterrupt, SystemExit) as exc: + state.fatal_error = exc + state._disconnect(1006) + state._cancel_handler() + raise + except BaseException: + if state.shutdown: + return + state._enqueue_control( + state.api.CloseConnection(code=1002, reason="protocol error"), 16 + ) + state.close_sent = True + state._disconnect(1002, "protocol error") + state._cancel_handler() + + +def _receive_protocol_data(state: _WebSocketState, data: bytes) -> None: + for chunk in state._guard.feed(data): + state.protocol.receive_data(chunk) + + +def _drain_protocol_events(state: _WebSocketState) -> bool: + for event in state.protocol.events(): + if isinstance(event, (state.api.TextMessage, state.api.BytesMessage)): + kind = str if isinstance(event, state.api.TextMessage) else bytes + if state._message_kind is None: + state._message_kind = kind + state._message_buffer.clear() + state._message_bytes = 0 + if state._message_kind is not kind: + raise ValueError("WebSocket message type changed during fragmentation") + state._message_bytes += _message_size(event.data) + if state._message_bytes > state.config.max_message_bytes: + raise WebSocketCapacityError("WebSocket message is too large") + encoded = ( + event.data.encode("utf-8") + if isinstance(event.data, str) + else event.data + ) + state._message_buffer.extend(encoded) + if event.message_finished: + value = ( + bytes(state._message_buffer).decode("utf-8") + if kind is str + else bytes(state._message_buffer) + ) + state._message_kind = None + state._message_buffer.clear() + state._message_bytes = 0 + if not state._deliver_message(value): + raise WebSocketCapacityError("WebSocket inbound queue is full") + elif isinstance(event, state.api.Ping): + if not state._enqueue_control(event.response(), len(event.payload)): + raise WebSocketCapacityError("WebSocket outbound queue is full") + elif isinstance(event, state.api.Pong): + _handle_pong(state, event.payload) + elif isinstance(event, state.api.CloseConnection): + state.peer_closed = True + close_code = int(event.code) + disconnect_reason = event.reason or "" + if not state.close_sent: + if close_code == 1002: + response = state.api.CloseConnection( + code=1002, reason="protocol error" + ) + elif close_code == 1007: + response = state.api.CloseConnection( + code=1007, reason="invalid payload" + ) + else: + response = event.response() + if not state._enqueue_control( + response, _close_event_size(response) + ): + raise WebSocketCapacityError("WebSocket outbound queue is full") + state.close_sent = True + if close_code == 1002: + disconnect_reason = "protocol error" + elif close_code == 1007: + disconnect_reason = "invalid payload" + state._disconnect(close_code, disconnect_reason) + return True + return False + + +def _handle_pong(state: _WebSocketState, payload: bytes) -> None: + if ( + state._pending_ping_generation is not None + and payload == state._pending_ping_payload + ): + state._clear_pending_ping() + + +def _close_event_size(event: Any) -> int: + if int(event.code) == 1005: + return 0 + return 2 + len((event.reason or "").encode("utf-8")) + + +async def _run_deadlines(task: Any, state: _WebSocketState) -> None: + while not state.shutdown: + await task.sleep(state.config.deadline_resolution) + now = time.monotonic() + if state.handshake_state in {"pending", "accepting", "rejecting"}: + if now - state.created_at >= state.config.handshake_timeout: + if state.handshake_state == "pending": + try: + await state.reject( + Response.text("WebSocket handshake timed out", status=408) + ) + except (KeyboardInterrupt, SystemExit) as exc: + state.fatal_error = exc + raise + except BaseException: + state._disconnect(1006) + else: + error = WebSocketDisconnect(1006, "handshake timed out") + state._fail_handshake(error) + state._cancel_handler() + return + continue + if state.pong_deadline is not None and now >= state.pong_deadline: + _begin_deadline_close(state, 1002, "Pong timeout") + elif state.accepted and now - state.last_activity >= state.config.idle_timeout: + _begin_deadline_close(state, 1001, "idle timeout") + if state.write_deadline is not None and now >= state.write_deadline: + error = WebSocketDisconnect(1006, "write timed out") + state._abort_writer(error) + state._cancel_handler() + state._disconnect(error.code, error.reason) + return + if state.close_deadline is not None and now >= state.close_deadline: + error = state.disconnect or WebSocketDisconnect(1006, "close timed out") + state._abort_writer(error) + state._cancel_handler() + state._disconnect(error.code, error.reason) + return + + +def _begin_deadline_close( + state: _WebSocketState, code: int, reason: str +) -> None: + if not state.close_sent and state.protocol is not None: + reason_bytes = reason.encode("utf-8") + state._enqueue_control( + state.api.CloseConnection(code=code, reason=reason), + 2 + len(reason_bytes), + ) + state.close_sent = True + if state.close_deadline is None: + state.close_deadline = time.monotonic() + state.config.close_timeout + state._disconnect(code, reason) + state._cancel_handler() diff --git a/tests/installed_regex_smoke.py b/tests/installed_regex_smoke.py new file mode 100644 index 0000000..e439c6d --- /dev/null +++ b/tests/installed_regex_smoke.py @@ -0,0 +1,23 @@ +"""Smoke an installed ``smallserver[regex-routes]`` package outside the source path.""" + +import asyncio + +from smallserver import Headers, Request, Response, RouteErrorEvent, SmallServer + + +async def main() -> None: + event = RouteErrorEvent("regex-route-smoke", "route_match_timeout") + assert (event.route_id, event.category) == ("regex-route-smoke", "route_match_timeout") + app = SmallServer() + + @app.get_regex(r"/users/(?P[0-9]+)") + async def user(request: Request) -> Response: + return Response.text(request.path_params["user_id"]) + + response = await app.dispatch(Request("GET", "/users/42?source=smoke", Headers())) + assert response.status == 200 + assert response.body == b"42" + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/tests/test_http.py b/tests/test_http.py index 16a9b9a..f98519d 100644 --- a/tests/test_http.py +++ b/tests/test_http.py @@ -1,4 +1,5 @@ import unittest +from types import MappingProxyType from smallserver import Headers, Request, Response @@ -38,3 +39,28 @@ def test_header_values_match_the_wire_encoding(self) -> None: with self.subTest(invalid=invalid): with self.assertRaisesRegex(ValueError, "header value"): Headers({"X-Test": invalid}) + + def test_request_splits_raw_target_and_keeps_captures_immutable(self) -> None: + request = Request( + "GET", + "/items?tag=a%2Fb&empty=", + Headers(), + path_params={"item_id": "a%2Fb"}, + route_pattern=r"/items/(?P[^/]+)", + ) + self.assertEqual(request.raw_target, "/items?tag=a%2Fb&empty=") + self.assertEqual(request.path, "/items") + self.assertEqual(request.query_string, "tag=a%2Fb&empty=") + self.assertIsInstance(request.path_params, MappingProxyType) + with self.assertRaises(TypeError): + request.path_params["item_id"] = "changed" # type: ignore[index] + + def test_request_rejects_inconsistent_explicit_target_fields(self) -> None: + with self.assertRaisesRegex(ValueError, "inconsistent"): + Request("GET", "/one", Headers(), raw_target="/two") + + def test_request_can_be_built_from_separate_path_and_query(self) -> None: + request = Request("GET", "/items", Headers(), query_string="page=2") + self.assertEqual(request.raw_target, "/items?page=2") + self.assertEqual(request.path, "/items") + self.assertEqual(request.query_string, "page=2") diff --git a/tests/test_lifecycle.py b/tests/test_lifecycle.py index 3f3648c..bda5989 100644 --- a/tests/test_lifecycle.py +++ b/tests/test_lifecycle.py @@ -130,6 +130,17 @@ def test_managed_priorities_are_validated_before_runtime_creation(self) -> None: SmallServer().listen(config=config, port=0) factory.assert_not_called() + observed = ServerConfig( + max_connections=32, + managed_runtime=ManagedRuntimeConfig(task_capacity=34), + ) + with patch("smallserver.app._default_runtime_factory") as factory: + with self.assertRaisesRegex(ValueError, r"max_connections \+ 3"): + SmallServer(route_error_observer=lambda event: None).listen( + config=observed, port=0 + ) + factory.assert_not_called() + insufficient = ServerConfig( max_connections=32, managed_runtime=ManagedRuntimeConfig(task_capacity=33), @@ -170,7 +181,11 @@ def test_primary_demo_hides_runtime_and_registers_all_http_methods(self) -> None import demo self.assertEqual( - {method for method, path in demo.app._routes if path == "/tasks"}, + { + method + for method, path in demo.app._router._static + if path == "/tasks" + }, {"GET", "POST", "PUT", "PATCH", "DELETE"}, ) diff --git a/tests/test_regex_routing.py b/tests/test_regex_routing.py new file mode 100644 index 0000000..71c0b73 --- /dev/null +++ b/tests/test_regex_routing.py @@ -0,0 +1,310 @@ +import importlib.util +import time +import unittest +from unittest.mock import patch + +from smallserver import ( + Headers, + RegexRouteConfig, + RegexRoutesUnavailable, + Request, + Response, + RouteMatchTimeout, + SmallServer, +) +from smallserver.routing import Router + + +HAS_REGEX = importlib.util.find_spec("regex") is not None + + +class OptionalRegexDependencyTests(unittest.TestCase): + def test_error_observer_must_be_callable(self) -> None: + with self.assertRaisesRegex(TypeError, "route_error_observer"): + SmallServer(route_error_observer=object()) # type: ignore[arg-type] + + def test_static_routes_do_not_import_optional_engine(self) -> None: + app = SmallServer() + + async def handler(request): + return Response() + + with patch("smallserver.routing.importlib.import_module", side_effect=AssertionError("imported")): + app.get("/health")(handler) + + def test_registration_explains_missing_extra(self) -> None: + app = SmallServer() + + async def handler(request): + return Response() + + with patch("smallserver.routing.importlib.import_module", side_effect=ImportError): + with self.assertRaisesRegex(RegexRoutesUnavailable, r"smallserver\[regex-routes\]"): + app.get_regex(r"/users/[0-9]+")(handler) + + +@unittest.skipUnless(HAS_REGEX, "regex-routes extra is not installed") +class RegexRoutingTests(unittest.IsolatedAsyncioTestCase): + async def test_named_captures_are_immutable_and_optional_groups_are_omitted(self) -> None: + app = SmallServer() + seen = [] + + @app.get_regex(r"/users/(?P[0-9]+)(?:/(?P
[a-z]+))?") + async def user(request): + seen.append(request) + return Response.json(dict(request.path_params)) + + response = await app.dispatch(Request("GET", "/users/42?debug=1", Headers())) + self.assertEqual(response.body, b'{"user_id":"42"}') + self.assertEqual(seen[0].raw_target, "/users/42?debug=1") + self.assertEqual(seen[0].query_string, "debug=1") + self.assertEqual(seen[0].route_pattern, r"/users/(?P[0-9]+)(?:/(?P
[a-z]+))?") + with self.assertRaises(TypeError): + seen[0].path_params["user_id"] = "1" # type: ignore[index] + await app.dispatch(Request("GET", "/users/43/profile", Headers())) + self.assertEqual(dict(seen[0].path_params), {"user_id": "42"}) + self.assertEqual(dict(seen[1].path_params), {"user_id": "43", "section": "profile"}) + + async def test_captures_preserve_percent_encoded_octets(self) -> None: + app = SmallServer() + + @app.get_regex(r"/files/(?P[^/]+)") + async def file(request): + return Response.text(request.path_params["name"]) + + response = await app.dispatch(Request("GET", "/files/a%2Fb", Headers())) + self.assertEqual(response.body, b"a%2Fb") + + async def test_static_precedence_and_regex_registration_order(self) -> None: + app = SmallServer() + + @app.get_regex(r"/items/(?P.+)") + async def broad(request): + return Response.text("broad") + + @app.get_regex(r"/items/(?P[0-9]+)") + async def narrow(request): + return Response.text("narrow") + + @app.get("/items/7") + async def exact(request): + return Response.text("static") + + self.assertEqual((await app.dispatch(Request("GET", "/items/7", Headers()))).body, b"static") + self.assertEqual((await app.dispatch(Request("GET", "/items/8", Headers()))).body, b"broad") + + async def test_alternation_and_optional_literals_preserve_order_and_405(self) -> None: + app = SmallServer() + + @app.get_regex(r"/foo|/bar") + async def top_level(request): + return Response.text("top") + + @app.get_regex(r"/fo?bar") + async def optional(request): + return Response.text("optional") + + @app.get_regex(r"/(?:red|blue)") + async def nested(request): + return Response.text("nested") + + for path, expected in ( + ("/foo", b"top"), + ("/bar", b"top"), + ("/fbar", b"optional"), + ("/fobar", b"optional"), + ("/red", b"nested"), + ("/blue", b"nested"), + ): + with self.subTest(path=path): + self.assertEqual((await app.dispatch(Request("GET", path, Headers()))).body, expected) + response = await app.dispatch(Request("POST", "/bar", Headers())) + self.assertEqual(response.status, 405) + self.assertEqual(response.headers["allow"], "GET") + + def test_compiled_routes_are_structurally_guarded_to_slash_paths(self) -> None: + async def handler(request): + return Response() + + for pattern, slash_path, non_slash_path in ( + (r"/foo|bar", "/foo", "bar"), + (r"/?foo", "/foo", "foo"), + ): + with self.subTest(pattern=pattern): + router = Router() + router.add_regex(pattern, ("GET",), handler) + self.assertIs(router.resolve("GET", slash_path).handler, handler) + miss = router.resolve("GET", non_slash_path) + self.assertIsNone(miss.handler) + self.assertEqual(miss.allowed_methods, ()) + + async def test_method_first_matching_and_sorted_allow(self) -> None: + app = SmallServer() + + @app.post("/records/1") + async def static_post(request): + return Response.text("post") + + @app.delete_regex(r"/records/(?P[0-9]+)") + async def regex_delete(request): + return Response.text("delete") + + @app.get_regex(r"/records/(?P.+)") + async def regex_get(request): + return Response.text("get") + + self.assertEqual((await app.dispatch(Request("GET", "/records/1", Headers()))).body, b"get") + response = await app.dispatch(Request("PATCH", "/records/1", Headers())) + self.assertEqual(response.status, 405) + self.assertEqual(response.headers["allow"], "DELETE, GET, POST") + + async def test_all_regex_method_decorators_dispatch(self) -> None: + app = SmallServer() + for method in ("get", "post", "put", "patch", "delete"): + decorator = getattr(app, method + "_regex") + + @decorator("/" + method + r"/(?P[0-9]+)") + async def handler(request, expected=method): + return Response.text(expected + request.path_params["value"]) + + for method in ("get", "post", "put", "patch", "delete"): + response = await app.dispatch(Request(method.upper(), "/" + method + "/2", Headers())) + self.assertEqual(response.body, (method + "2").encode()) + + async def test_same_pattern_disjoint_methods_merge_and_duplicates_are_atomic(self) -> None: + app = SmallServer() + + async def first(request): + return Response.text("first") + + async def second(request): + return Response.text("second") + + app.get_regex(r"/merged/(?P[0-9]+)")(first) + app.post_regex(r"/merged/(?P[0-9]+)")(second) + with self.assertRaisesRegex(ValueError, "already registered"): + app.route_regex(r"/merged/(?P[0-9]+)", ("PATCH", "GET"))(second) + self.assertEqual((await app.dispatch(Request("PATCH", "/merged/1", Headers()))).status, 405) + self.assertEqual((await app.dispatch(Request("POST", "/merged/1", Headers()))).body, b"second") + + def test_registration_limits_and_invalid_patterns(self) -> None: + async def handler(request): + return Response() + + cases = ( + (r"items/[0-9]+", "literal '/'"), + (r"/[", "invalid"), + (r"/(?Px)(?Py)", "duplicate"), + ) + for pattern, message in cases: + with self.subTest(pattern=pattern): + with self.assertRaisesRegex((ValueError, TypeError), message): + SmallServer().get_regex(pattern)(handler) + + with self.assertRaisesRegex(ValueError, "too long"): + SmallServer(RegexRouteConfig(max_pattern_length=4)).get_regex("/long")(handler) + with self.assertRaisesRegex(ValueError, "too many"): + SmallServer(RegexRouteConfig(max_named_captures=1)).get_regex( + r"/(?Px)(?Py)" + )(handler) + limited = SmallServer(RegexRouteConfig(max_routes=1)) + limited.get_regex(r"/one")(handler) + with self.assertRaisesRegex(ValueError, "maximum"): + limited.get_regex(r"/two")(handler) + + async def test_verbose_comment_group_text_does_not_create_false_duplicate(self) -> None: + app = SmallServer() + + @app.get_regex( + "/(?x:items/ # (?Pthis is comment text)\n" + " (?P[0-9]+))" + ) + async def item(request): + return Response.text(request.path_params["item_id"]) + + response = await app.dispatch(Request("GET", "/items/42", Headers())) + self.assertEqual(response.body, b"42") + + async def test_catastrophic_backtracking_is_bounded_and_path_is_not_disclosed(self) -> None: + app = SmallServer(RegexRouteConfig(match_timeout=0.001, total_match_timeout=0.005)) + + @app.get_regex(r"/(a+)+$") + async def handler(request): + return Response() + + hostile_path = "/" + "a" * 5000 + "!" + started = time.monotonic() + with self.assertRaises(RouteMatchTimeout) as raised: + await app.dispatch(Request("GET", hostile_path, Headers())) + elapsed = time.monotonic() - started + self.assertLess(elapsed, 0.25) + self.assertNotIn(hostile_path, str(raised.exception)) + self.assertEqual(raised.exception.route_id, "regex-route-1") + + async def test_total_budget_bounds_many_individually_fast_misses(self) -> None: + class SlowPattern: + groupindex = {} + + def fullmatch(self, value, timeout): + time.sleep(min(timeout / 2, 0.001)) + return None + + class Engine: + @staticmethod + def compile(pattern): + return SlowPattern() + + app = SmallServer( + RegexRouteConfig(max_routes=20, match_timeout=0.01, total_match_timeout=0.003) + ) + + async def handler(request): + return Response() + + with patch("smallserver.routing.importlib.import_module", return_value=Engine()): + for index in range(10): + app.get_regex("/(?:route-{}).*".format(index))(handler) + started = time.monotonic() + with self.assertRaises(RouteMatchTimeout): + await app.dispatch(Request("GET", "/route", Headers())) + self.assertLess(time.monotonic() - started, 0.05) + + async def test_manual_dispatch_path_limit_is_enforced_before_matching(self) -> None: + calls = [] + + class Pattern: + groupindex = {} + + def fullmatch(self, value, timeout): + calls.append(value) + return None + + class Engine: + @staticmethod + def compile(pattern): + return Pattern() + + app = SmallServer(RegexRouteConfig(max_path_bytes=8)) + + async def handler(request): + return Response() + + with patch("smallserver.routing.importlib.import_module", return_value=Engine()): + app.get_regex(r"/.*")(handler) + response = await app.dispatch(Request("GET", "/12345678", Headers())) + self.assertEqual(response.status, 414) + self.assertEqual(calls, []) + + +class RegexRouteConfigTests(unittest.TestCase): + def test_rejects_nonfinite_or_nonpositive_limits(self) -> None: + for kwargs in ( + {"max_routes": 0}, + {"max_routes": True}, + {"match_timeout": 0}, + {"match_timeout": float("inf")}, + {"total_match_timeout": float("nan")}, + ): + with self.subTest(kwargs=kwargs): + with self.assertRaises(ValueError): + RegexRouteConfig(**kwargs) diff --git a/tests/test_route_benchmark.py b/tests/test_route_benchmark.py new file mode 100644 index 0000000..8b0abdd --- /dev/null +++ b/tests/test_route_benchmark.py @@ -0,0 +1,35 @@ +import contextlib +import io +import sys +import unittest +from unittest.mock import patch + +from benchmarks import route_benchmark + + +class RouteBenchmarkCLITests(unittest.TestCase): + def test_release_mode_rejects_undersized_samples(self) -> None: + cases = ( + ("9999", "5", "10000 iterations"), + ("10000", "4", "5 rounds"), + ) + for iterations, rounds, message in cases: + with self.subTest(iterations=iterations, rounds=rounds): + stderr = io.StringIO() + with patch.object( + sys, + "argv", + [ + "route_benchmark.py", + "--release", + "--iterations", + iterations, + "--rounds", + rounds, + ], + ): + with contextlib.redirect_stderr(stderr): + with self.assertRaises(SystemExit) as raised: + route_benchmark.main() + self.assertEqual(raised.exception.code, 2) + self.assertIn(message, stderr.getvalue()) diff --git a/tests/test_routing.py b/tests/test_routing.py index 19da695..e6aaed4 100644 --- a/tests/test_routing.py +++ b/tests/test_routing.py @@ -30,6 +30,58 @@ async def items(request): self.assertEqual(response.status, 405) self.assertEqual(response.headers["allow"], "GET") + async def test_query_string_does_not_participate_in_static_matching(self) -> None: + app = SmallServer() + seen = [] + + @app.get("/items") + async def items(request): + seen.append(request) + return Response() + + response = await app.dispatch(Request("GET", "/items?tag=a%2Fb", Headers())) + self.assertEqual(response.status, 200) + self.assertEqual(seen[0].path, "/items") + self.assertEqual(seen[0].query_string, "tag=a%2Fb") + + async def test_static_dispatch_passes_original_request_without_route_context_copy(self) -> None: + app = SmallServer() + seen = [] + + @app.get("/health") + async def health(request): + seen.append(request) + return Response() + + request = Request("GET", "/health", Headers()) + response = await app.dispatch(request) + self.assertEqual(response.status, 200) + self.assertIs(seen[0], request) + self.assertIsNone(seen[0].route_pattern) + self.assertEqual(dict(seen[0].path_params), {}) + + async def test_static_dispatch_clears_spoofed_route_context(self) -> None: + app = SmallServer() + seen = [] + + @app.get("/health") + async def health(request): + seen.append(request) + return Response() + + dirty = Request( + "GET", + "/health", + Headers(), + path_params={"spoofed": "value"}, + route_pattern="sensitive-spoofed-pattern", + ) + response = await app.dispatch(dirty) + self.assertEqual(response.status, 200) + self.assertIsNot(seen[0], dirty) + self.assertIsNone(seen[0].route_pattern) + self.assertEqual(dict(seen[0].path_params), {}) + async def test_http_error_becomes_response(self) -> None: app = SmallServer() @@ -61,6 +113,8 @@ async def one(request): app.get("/one")(one) with self.assertRaisesRegex(ValueError, "supported HTTP methods"): app.route("/trace", ("TRACE",)) + with self.assertRaisesRegex(ValueError, "query string"): + app.get("/one?debug=1") async def test_failed_multi_method_registration_is_atomic(self) -> None: app = SmallServer() diff --git a/tests/test_server.py b/tests/test_server.py index 8c5938a..57e0f21 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -4,9 +4,19 @@ import warnings from unittest.mock import patch -from smallserver import ManagedRuntimeConfig, ServerStartupError, SmallServer +from smallserver import ( + ManagedRuntimeConfig, + RouteErrorEvent, + ServerStartupError, + SmallServer, +) from smallserver.errors import _CleanupTransaction -from smallserver.server import HTTPParseError, HTTPRequestParser, ServerConfig +from smallserver.server import ( + HTTPParseError, + HTTPRequestParser, + RouteObserverChannel, + ServerConfig, +) from tests.kernel_fakes import FakeKernel @@ -41,14 +51,85 @@ def test_requires_host_and_rejects_invalid_origin_form(self) -> None: with self.assertRaisesRegex(HTTPParseError, "origin-form"): self.parser().feed(b"GET /items#fragment HTTP/1.1\r\nHost: localhost\r\n\r\n") + def test_splits_query_without_decoding_or_normalizing_path(self) -> None: + request = self.parser().feed( + b"GET /items/a%2Fb?tag=x%20y HTTP/1.1\r\nHost: localhost\r\n\r\n" + ) + self.assertIsNotNone(request) + assert request is not None + self.assertEqual(request.raw_target, "/items/a%2Fb?tag=x%20y") + self.assertEqual(request.path, "/items/a%2Fb") + self.assertEqual(request.query_string, "tag=x%20y") + + def test_enforces_request_target_limit_independently(self) -> None: + parser = HTTPRequestParser(256, 2, 32, max_request_target_bytes=8) + with self.assertRaisesRegex(HTTPParseError, "request target") as raised: + parser.feed(b"GET /12345678 HTTP/1.1\r\nHost: x\r\n\r\n") + self.assertEqual(raised.exception.status, 414) + def test_config_rejects_unbounded_limits(self) -> None: with self.assertRaisesRegex(ValueError, "max_connections"): ServerConfig(max_connections=0) with self.assertRaisesRegex(ValueError, "max_connections"): ServerConfig(max_connections=True) + with self.assertRaisesRegex(ValueError, "max_request_target_bytes"): + ServerConfig(max_request_target_bytes=0) + with self.assertRaisesRegex(ValueError, "max_route_error_events"): + ServerConfig(max_route_error_events=0) with self.assertRaisesRegex(TypeError, "managed_runtime"): ServerConfig(managed_runtime={}) # type: ignore[arg-type] + def test_config_preserves_legacy_positional_field_mapping(self) -> None: + config = ServerConfig(1, 2, 3, 4, 5, 6, 7) + self.assertEqual(config.max_connections, 1) + self.assertEqual(config.max_header_bytes, 2) + self.assertEqual(config.max_header_count, 3) + self.assertEqual(config.max_body_bytes, 4) + self.assertEqual(config.receive_chunk_bytes, 5) + self.assertEqual(config.listener_priority, 6) + self.assertEqual(config.connection_priority, 7) + self.assertEqual(config.max_request_target_bytes, 8 * 1024) + self.assertEqual(config.max_route_error_events, 16) + + def test_route_observer_channel_has_deterministic_capacity_and_stop(self) -> None: + class ObserverTask: + @staticmethod + def getID() -> int: + return 9 + + class SourceTask: + signals = [] + + def sendSignal(self, pid, signal) -> int: + self.signals.append((pid, signal)) + return 0 + + channel = RouteObserverChannel(lambda event: None, max_events=1) + channel.bind(ObserverTask()) + source = SourceTask() + first = RouteErrorEvent("regex-route-1", "route_match_timeout") + second = RouteErrorEvent("regex-route-2", "route_match_timeout") + self.assertTrue(channel.enqueue(first, source)) + self.assertFalse(channel.enqueue(second, source)) + self.assertEqual(list(channel.events), [first]) + self.assertEqual(channel.dropped, 1) + self.assertEqual(source.signals, [(9, 31)]) + channel.stop() + self.assertFalse(channel.accepting) + self.assertEqual(list(channel.events), []) + self.assertEqual(channel.dropped, 2) + + failing_channel = RouteObserverChannel(lambda event: None, max_events=1) + failing_channel.bind(ObserverTask()) + + class FailingSourceTask: + def sendSignal(self, pid, signal) -> int: + raise RuntimeError("signal failed") + + self.assertFalse(failing_channel.enqueue(first, FailingSourceTask())) + self.assertEqual(list(failing_channel.events), []) + self.assertEqual(failing_channel.dropped, 1) + def test_managed_runtime_config_is_validated_and_defensively_copied(self) -> None: source = {"http": {"max_response_size": 4096}} config = ManagedRuntimeConfig( @@ -107,6 +188,39 @@ def resume_task(self, task) -> None: self.assertEqual([handle.name for handle in runtime.kernel.closed], ["listener"]) self.assertEqual(runtime.kernel.wakeup.close_calls, 1) + def test_observer_task_is_owned_by_startup_rollback(self) -> None: + class Runtime: + def __init__(self) -> None: + self.kernel = FakeKernel() + self.tasks = [] + self.cancelled = [] + + def fork(self, tasks) -> None: + self.tasks = list(tasks) + raise RuntimeError("no task capacity") + + def cancel_task(self, task) -> None: + self.cancelled.append(task) + task.cancel() + + def resume_task(self, task) -> None: + pass + + runtime = Runtime() + app = SmallServer(route_error_observer=lambda event: None) + with self.assertRaisesRegex(RuntimeError, "capacity"): + app.serve(runtime) + + self.assertEqual(runtime.cancelled, runtime.tasks) + self.assertEqual( + [task.name for task in runtime.tasks], + [ + "smallserver-listener", + "smallserver-close-watcher", + "smallserver-route-observer", + ], + ) + def test_serve_closes_kernel_resources_when_task_construction_fails(self) -> None: from SmallPackage import SmallTask as RealSmallTask diff --git a/tests/test_server_runtime.py b/tests/test_server_runtime.py index 07d8e59..18441cf 100644 --- a/tests/test_server_runtime.py +++ b/tests/test_server_runtime.py @@ -1,3 +1,6 @@ +import importlib.util +from dataclasses import FrozenInstanceError +import inspect import socket import threading import time @@ -6,16 +9,25 @@ from SmallPackage import SmallOS, Unix from SmallPackage.adapters.threads import ThreadAdapter -from smallserver import AdapterRegistry, Response, SmallServer -from smallserver.server import ServerHandle +from smallserver import ( + AdapterRegistry, + RegexRouteConfig, + Request, + Response, + RouteErrorEvent, + RouteMatchTimeout, + SmallServer, +) +from smallserver.server import ServerHandle, run_route_observer + + +HAS_REGEX = importlib.util.find_spec("regex") is not None class SmallOSServerIntegrationTests(unittest.TestCase): - def _request(self, port: int, path: str) -> bytes: + def _exchange(self, port: int, payload: bytes) -> bytes: with socket.create_connection(("127.0.0.1", port), timeout=3) as connection: - connection.sendall( - "GET {} HTTP/1.1\r\nHost: localhost\r\n\r\n".format(path).encode("ascii") - ) + connection.sendall(payload) chunks = [] while True: chunk = connection.recv(4096) @@ -23,6 +35,12 @@ def _request(self, port: int, path: str) -> bytes: return b"".join(chunks) chunks.append(chunk) + def _request(self, port: int, path: str) -> bytes: + return self._exchange( + port, + "GET {} HTTP/1.1\r\nHost: localhost\r\n\r\n".format(path).encode("ascii"), + ) + def test_loopback_server_accepts_fragmented_request_and_shuts_down(self) -> None: runtime = SmallOS().setKernel(Unix()) app = SmallServer() @@ -136,6 +154,180 @@ def fast_client() -> None: self.assertIn(b"\r\n\r\nfast done", responses["fast"]) self.assertIn(b"\r\n\r\nslow done", responses["slow"]) + @unittest.skipUnless(HAS_REGEX, "regex-routes extra is not installed") + def test_loopback_regex_route_uses_path_without_query(self) -> None: + runtime = SmallOS().setKernel(Unix()) + app = SmallServer() + + @app.get_regex(r"/files/(?P[^/]+)") + async def file(request): + return Response.text(request.path_params["name"] + "?" + request.query_string) + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + received = [] + errors = [] + + def client() -> None: + try: + received.append(self._request(server.port, "/files/a%2Fb?download=1")) + except BaseException as exc: + errors.append(exc) + finally: + server.close() + + worker = threading.Thread(target=client, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=2) + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertIn(b"\r\n\r\na%2Fb?download=1", b"".join(received)) + + @unittest.skipUnless(HAS_REGEX, "regex-routes extra is not installed") + def test_regex_timeout_is_observed_once_and_does_not_stop_server(self) -> None: + runtime = SmallOS().setKernel(Unix()) + observed: list[RouteErrorEvent] = [] + observer_graph = [] + observer_finished = threading.Event() + observer_threads = [] + + def observe(event: RouteErrorEvent) -> None: + observed.append(event) + observer_threads.append(threading.current_thread()) + caller_locals = [] + frame = inspect.currentframe() + while frame is not None: + caller_locals.append(dict(frame.f_locals)) + if frame.f_code is run_route_observer.__code__: + break + frame = frame.f_back + observer_graph.extend(_reachable_container_values(caller_locals)) + observer_finished.set() + raise RuntimeError("intentional observer failure") + + app = SmallServer( + RegexRouteConfig(match_timeout=0.001, total_match_timeout=0.005), + route_error_observer=observe, + ) + + pattern_secret = "sensitive-pattern-marker" + + @app.post_regex(r"/(a+)+$(?#sensitive-pattern-marker)") + async def expensive(request): + return Response() + + @app.get("/health") + async def health(request): + return Response.text("healthy") + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + hostile_path = "/" + "a" * 5000 + "!" + authorization_secret = "Bearer sensitive-authorization-marker" + body_secret = b"sensitive-body-marker" + runtime_thread = threading.Thread( + target=runtime.start, + name="smallos-runtime-test", + daemon=True, + ) + runtime_thread.start() + request = ( + "POST {} HTTP/1.1\r\n" + "Host: localhost\r\n" + "Authorization: {}\r\n" + "Content-Length: {}\r\n\r\n" + ).format(hostile_path, authorization_secret, len(body_secret)).encode("ascii") + received = [self._exchange(server.port, request + body_secret)] + received.append(self._request(server.port, "/health")) + self.assertTrue(observer_finished.wait(2), "route observer did not run") + server.close() + runtime_thread.join(timeout=3) + self.assertFalse(runtime_thread.is_alive()) + self.assertEqual(len(observed), 1) + event = observed[0] + self.assertEqual(event.route_id, "regex-route-1") + self.assertEqual(event.category, "route_match_timeout") + with self.assertRaises(FrozenInstanceError): + event.route_id = "changed" # type: ignore[misc] + self.assertFalse(hasattr(event, "__traceback__")) + self.assertFalse(hasattr(event, "__cause__")) + self.assertFalse(hasattr(event, "__context__")) + + reachable = _reachable_objects(event) + reachable_strings = {value for value in reachable if isinstance(value, str)} + self.assertEqual( + reachable_strings, + {"route_id", "category", "regex-route-1", "route_match_timeout"}, + ) + self.assertFalse(any(isinstance(value, Request) for value in reachable)) + for secret in (hostile_path, authorization_secret, body_secret.decode("ascii"), pattern_secret): + self.assertNotIn(secret, reachable_strings) + + caller_strings = {value for value in observer_graph if isinstance(value, str)} + self.assertFalse(any(isinstance(value, Request) for value in observer_graph)) + self.assertFalse(any(isinstance(value, RouteMatchTimeout) for value in observer_graph)) + for secret in (hostile_path, authorization_secret, body_secret.decode("ascii"), pattern_secret): + self.assertNotIn(secret, caller_strings) + self.assertNotIn(body_secret, observer_graph) + self.assertEqual(server.route_observer_failures, 1) + self.assertEqual(server.dropped_route_error_events, 0) + self.assertEqual(observer_threads, [runtime_thread]) + self.assertNotIn( + "smallserver-route-observer", + {thread.name for thread in threading.enumerate()}, + ) + channel = server._route_observer_channel + self.assertIsNotNone(channel) + assert channel is not None + self.assertFalse(channel.accepting) + self.assertEqual(list(channel.events), []) + self.assertIsNone(channel.task) + self.assertTrue(received[0].startswith(b"HTTP/1.1 500 Internal Server Error\r\n")) + self.assertNotIn(hostile_path.encode("ascii"), received[0]) + self.assertTrue(received[1].startswith(b"HTTP/1.1 200 OK\r\n")) + self.assertTrue(received[1].endswith(b"healthy")) + + @unittest.skipUnless(HAS_REGEX, "regex-routes extra is not installed") + def test_regex_path_limit_returns_414_before_matching(self) -> None: + runtime = SmallOS().setKernel(Unix()) + app = SmallServer(RegexRouteConfig(max_path_bytes=8)) + + @app.get_regex(r"/.*") + async def route(request): + return Response.text("must not run") + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + received = [] + errors = [] + + def client() -> None: + try: + received.append(self._request(server.port, "/12345678")) + except BaseException as exc: + errors.append(exc) + finally: + server.close() + + worker = threading.Thread(target=client, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=2) + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertTrue(received[0].startswith(b"HTTP/1.1 414 URI Too Long\r\n")) + self.assertNotIn(b"must not run", received[0]) + def test_slow_client_does_not_block_a_complete_request(self) -> None: runtime = SmallOS().setKernel(Unix()) app = SmallServer() @@ -157,8 +349,12 @@ def clients() -> None: try: slow = socket.create_connection(("127.0.0.1", server.port), timeout=2) slow.sendall(b"GET /slow HTTP/1.1\r\nHost: local") - with socket.create_connection(("127.0.0.1", server.port), timeout=2) as fast_client: - fast_client.sendall(b"GET /fast HTTP/1.1\r\nHost: localhost\r\n\r\n") + with socket.create_connection( + ("127.0.0.1", server.port), timeout=2 + ) as fast_client: + fast_client.sendall( + b"GET /fast HTTP/1.1\r\nHost: localhost\r\n\r\n" + ) while True: chunk = fast_client.recv(4096) if not chunk: @@ -229,3 +425,43 @@ def run_server() -> None: self.assertTrue(handle.closed) self.assertEqual(handle.port, returned[0].port) self.assertIn(b"HTTP/1.1 200 OK", response) + + +def _reachable_objects(root): + pending = [root] + seen = set() + result = [] + while pending: + value = pending.pop() + identity = id(value) + if identity in seen: + continue + seen.add(identity) + result.append(value) + if isinstance(value, dict): + pending.extend(value.keys()) + pending.extend(value.values()) + elif isinstance(value, (tuple, list, set, frozenset)): + pending.extend(value) + elif hasattr(value, "__dict__"): + pending.append(vars(value)) + return result + +def _reachable_container_values(root): + """Walk frame-local containers without traversing scheduler object graphs.""" + pending = [root] + seen = set() + result = [] + while pending: + value = pending.pop() + identity = id(value) + if identity in seen: + continue + seen.add(identity) + result.append(value) + if isinstance(value, dict): + pending.extend(value.keys()) + pending.extend(value.values()) + elif isinstance(value, (tuple, list, set, frozenset)): + pending.extend(value) + return result diff --git a/tests/test_websocket.py b/tests/test_websocket.py new file mode 100644 index 0000000..9755560 --- /dev/null +++ b/tests/test_websocket.py @@ -0,0 +1,1174 @@ +from __future__ import annotations + +import base64 +import importlib.util +import socket +import threading +import time +import unittest +from unittest.mock import patch + +from SmallPackage import SmallOS, SmallTask, SmallWebSocketClient, Unix +from SmallPackage.adapters.threads import ThreadAdapter + +from smallserver import ( + AdapterRegistry, + Headers, + Request, + Response, + SmallServer, + WebSocket, + WebSocketCapacityError, + WebSocketConfig, + WebSocketDisconnect, + WebSocketUnavailable, +) +from smallserver.websocket import ( + _FrameGuard, + _WebSocketState, + _WebSocketRoute, + _drain_protocol_events, + _handle_pong, + _load_wsproto, + _receive_protocol_data, + _run_deadlines, + _run_handler, + _run_writer, + _validate_upgrade, +) + + +HAS_WSPROTO = importlib.util.find_spec("wsproto") is not None + + +def upgrade_request(**headers: str) -> Request: + values = { + "Host": "localhost", + "Upgrade": "websocket", + "Connection": "keep-alive, Upgrade", + "Sec-WebSocket-Version": "13", + "Sec-WebSocket-Key": "dGhlIHNhbXBsZSBub25jZQ==", + } + values.update(headers) + return Request("GET", "/ws", Headers(values)) + + +async def unused_handler(socket: WebSocket) -> None: + await socket.reject(Response(status=403)) + + +class _RecordingTransport: + def __init__(self) -> None: + self.send_calls = 0 + self.payloads: list[bytes] = [] + + async def send_all(self, task, client, payload: bytes) -> None: + self.send_calls += 1 + self.payloads.append(payload) + + +class _BlockingTransport(_RecordingTransport): + async def send_all(self, task, client, payload: bytes) -> None: + self.send_calls += 1 + self.payloads.append(payload) + await task.wait_signal(28) + + +def _make_state( + runtime, + transport, + handler, + *, + config: WebSocketConfig | None = None, +) -> _WebSocketState: + return _WebSocketState( + runtime, + transport, + object(), + upgrade_request(), + _WebSocketRoute(handler, None, ()), + config or WebSocketConfig(), + b"", + ) + + +def _make_accepted_state( + runtime, + transport, + handler, + *, + config: WebSocketConfig | None = None, +) -> _WebSocketState: + state = _make_state(runtime, transport, handler, config=config) + state.protocol = state.api.Connection(state.api.ConnectionType.SERVER) + state.accepted = True + state.handshake_state = "accepted" + return state + + +class WebSocketProtocolTests(unittest.TestCase): + def test_config_rejects_unbounded_or_inconsistent_limits(self) -> None: + with self.assertRaisesRegex(ValueError, "max_message_bytes"): + WebSocketConfig(max_message_bytes=0) + with self.assertRaisesRegex(ValueError, "must not exceed"): + WebSocketConfig(max_frame_payload_bytes=2, max_message_bytes=1) + with self.assertRaisesRegex(ValueError, "idle_timeout"): + WebSocketConfig(idle_timeout=float("inf")) + with self.assertRaisesRegex(ValueError, "write_timeout"): + WebSocketConfig(write_timeout=0) + + def test_registration_is_lazy_and_coexists_with_get(self) -> None: + app = SmallServer() + + @app.get("/ws") + async def ordinary(request): + return Response.text("http") + + with patch("importlib.import_module") as importer: + + @app.websocket("/ws") + async def websocket(socket): + await socket.accept() + + importer.assert_not_called() + self.assertIsNotNone(app._router.static_handler("GET", "/ws")) + self.assertIn("/ws", app._websocket_routes) + + def test_missing_optional_engine_has_clear_error(self) -> None: + real_import = __import__("importlib").import_module + + def missing(name: str): + if name.startswith("wsproto"): + raise ImportError("missing") + return real_import(name) + + with patch("importlib.import_module", side_effect=missing): + with self.assertRaisesRegex(WebSocketUnavailable, "websocket.*extra"): + _load_wsproto() + + def test_upgrade_validation_and_rfc_example_accept_inputs(self) -> None: + route = _WebSocketRoute(unused_handler, None, ("chat.v1",)) + self.assertIsNone(_validate_upgrade(upgrade_request(), route)) + response = _validate_upgrade( + upgrade_request(**{"Sec-WebSocket-Version": "12"}), route + ) + self.assertIsNotNone(response) + assert response is not None + self.assertEqual(response.status, 426) + self.assertEqual(response.headers["sec-websocket-version"], "13") + for name, value in ( + ("Upgrade", "not-websocket"), + ("Connection", "keep-alive"), + ("Sec-WebSocket-Key", base64.b64encode(b"short").decode("ascii")), + ): + with self.subTest(name=name): + self.assertEqual( + _validate_upgrade(upgrade_request(**{name: value}), route).status, + 400, + ) + + def test_origin_policy_is_explicit(self) -> None: + route = _WebSocketRoute( + unused_handler, frozenset({"https://allowed.example"}), () + ) + denied = _validate_upgrade( + upgrade_request(Origin="https://denied.example"), route + ) + self.assertIsNotNone(denied) + assert denied is not None + self.assertEqual(denied.status, 403) + self.assertIsNone( + _validate_upgrade( + upgrade_request(Origin="https://allowed.example"), route + ) + ) + + def test_subprotocol_tokens_are_ascii(self) -> None: + app = SmallServer() + for invalid in ("chat:v1", "café", "chat v1", ""): + with self.subTest(invalid=invalid): + with self.assertRaisesRegex(ValueError, "HTTP tokens"): + app.websocket( + "/invalid-{}".format(len(app._websocket_routes)), + subprotocols=(invalid,), + ) + + route = _WebSocketRoute(unused_handler, None, ("chat.v1",)) + for offered in ("chat:v1", "café", ","): + with self.subTest(offered=offered): + invalid = _validate_upgrade( + upgrade_request(**{"Sec-WebSocket-Protocol": offered}), route + ) + self.assertIsNotNone(invalid) + assert invalid is not None + self.assertEqual(invalid.status, 400) + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_extensions_cannot_be_selected_in_accept_response(self) -> None: + runtime = SmallOS().setKernel(Unix()) + transport = _RecordingTransport() + state = _make_state(runtime, transport, unused_handler) + outcome: list[BaseException] = [] + + async def accept_with_extension(task) -> None: + try: + await WebSocket(state).accept( + headers={"Sec-WebSocket-Extensions": "permessage-deflate"} + ) + except BaseException as exc: + outcome.append(exc) + + attempt = SmallTask(2, accept_with_extension, name="extension-rejection") + runtime.fork(attempt) + runtime.start() + self.assertIsInstance(outcome[0], ValueError) + self.assertEqual(state.handshake_state, "pending") + self.assertEqual(transport.send_calls, 0) + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_atomic_handshake_timeout_is_terminal_while_send_is_blocked(self) -> None: + runtime = SmallOS().setKernel(Unix()) + transport = _BlockingTransport() + config = WebSocketConfig( + handshake_timeout=0.02, + idle_timeout=1, + close_timeout=0.02, + deadline_resolution=0.005, + ) + state = _make_state(runtime, transport, unused_handler, config=config) + outcomes: list[BaseException] = [] + + async def accept_job(task) -> None: + state.handler_task = task + try: + await WebSocket(state).accept() + except BaseException as exc: + outcomes.append(exc) + + accept_task = SmallTask(2, accept_job, name="blocked-handshake") + deadline_task = SmallTask( + 2, _run_deadlines, args=(state,), name="handshake-deadline" + ) + state.deadline_task = deadline_task + runtime.fork([accept_task, deadline_task]) + runtime.start() + + self.assertEqual(transport.send_calls, 1) + self.assertEqual(state.handshake_state, "failed") + self.assertFalse(state.accepted) + self.assertFalse(state.rejected) + self.assertIsNotNone(state.handshake_error) + + async def retry_job(task) -> None: + try: + await WebSocket(state).accept() + except BaseException as exc: + outcomes.append(exc) + + retry_task = SmallTask(2, retry_job, name="handshake-retry") + runtime.fork(retry_task) + runtime.start() + self.assertIsInstance(outcomes[-1], Exception) + self.assertIn("already decided", str(outcomes[-1])) + self.assertEqual(transport.send_calls, 1) + + pending = _make_state( + SmallOS().setKernel(Unix()), _RecordingTransport(), unused_handler + ) + pending.request_shutdown() + self.assertEqual(pending.handshake_state, "failed") + self.assertIsInstance(pending.handshake_error, WebSocketDisconnect) + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_ping_is_armed_before_send_and_only_matching_pong_clears(self) -> None: + runtime = SmallOS().setKernel(Unix()) + state = _make_accepted_state(runtime, _RecordingTransport(), unused_handler) + observations: list[tuple[object, ...]] = [] + _handle_pong(state, b"unsolicited") + self.assertIsNone(state._pending_ping_generation) + + async def immediate_pong(event, size, *, wait): + observations.append( + ( + state._pending_ping_generation, + state._pending_ping_payload, + state.pong_deadline is not None, + ) + ) + _handle_pong(state, b"wrong") + observations.append((state._pending_ping_generation,)) + _handle_pong(state, b"probe") + + state._enqueue = immediate_pong + + async def ping_job(task) -> None: + await WebSocket(state).ping(b"probe") + + ping_task = SmallTask(2, ping_job, name="fast-pong") + runtime.fork(ping_task) + runtime.start() + self.assertIsNone(ping_task.exception) + self.assertEqual(observations[0][1:], (b"probe", True)) + self.assertIsNotNone(observations[1][0]) + self.assertIsNone(state._pending_ping_generation) + self.assertIsNone(state.pong_deadline) + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_fragment_metadata_is_coalesced_and_close_reasons_are_sanitized(self) -> None: + api = _load_wsproto() + runtime = SmallOS().setKernel(Unix()) + state = _make_accepted_state(runtime, _RecordingTransport(), unused_handler) + client = api.Connection(api.ConnectionType.CLIENT) + + _receive_protocol_data( + state, + client.send( + api.TextMessage(data="a", message_finished=False) + ), + ) + _drain_protocol_events(state) + for _ in range(2048): + _receive_protocol_data( + state, + client.send(api.TextMessage(data="", message_finished=False)), + ) + _drain_protocol_events(state) + self.assertEqual(len(state._message_buffer), 1) + _receive_protocol_data( + state, + client.send(api.TextMessage(data="b", message_finished=True)), + ) + _drain_protocol_events(state) + self.assertEqual(state.inbox.popleft().text, "ab") + + peer_close_state = _make_accepted_state( + runtime, _RecordingTransport(), unused_handler + ) + peer = api.Connection(api.ConnectionType.CLIENT) + _receive_protocol_data( + peer_close_state, + peer.send(api.CloseConnection(code=1000, reason="peer detail")), + ) + self.assertTrue(_drain_protocol_events(peer_close_state)) + command = peer_close_state.outbox.popleft() + self.assertEqual(command.size, 2 + len(b"peer detail")) + self.assertEqual(command.event.reason, "peer detail") + + protocol_error_state = _make_accepted_state( + runtime, _RecordingTransport(), unused_handler + ) + _receive_protocol_data(protocol_error_state, b"\x83\x80mask") + self.assertTrue(_drain_protocol_events(protocol_error_state)) + generated = protocol_error_state.outbox.popleft() + self.assertEqual(int(generated.event.code), 1002) + self.assertEqual(generated.event.reason, "protocol error") + self.assertEqual(protocol_error_state.disconnect.reason, "protocol error") + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_slow_writer_is_interrupted_by_bounded_deadlines(self) -> None: + runtime = SmallOS().setKernel(Unix()) + config = WebSocketConfig( + idle_timeout=1, + write_timeout=0.02, + close_timeout=0.02, + pong_timeout=1, + deadline_resolution=0.005, + ) + state = _make_accepted_state( + runtime, _BlockingTransport(), unused_handler, config=config + ) + outcomes: list[str] = [] + + async def sender(task) -> None: + state.handler_task = task + try: + await WebSocket(state).send_text("blocked") + finally: + outcomes.append("sender-finished") + + sender_task = SmallTask(2, sender, name="blocked-sender") + writer_task = SmallTask(2, _run_writer, args=(state,), name="blocked-writer") + deadline_task = SmallTask(2, _run_deadlines, args=(state,), name="write-deadline") + state.writer_task = writer_task + state.deadline_task = deadline_task + started = time.monotonic() + runtime.fork([sender_task, writer_task, deadline_task]) + runtime.start() + + self.assertLess(time.monotonic() - started, 0.5) + self.assertEqual(outcomes, ["sender-finished"]) + self.assertIsNotNone(state.disconnect) + self.assertEqual(state.disconnect.code, 1006) + self.assertEqual(state.disconnect.reason, "write timed out") + self.assertEqual(state.outbox_bytes, 0) + self.assertEqual(len(state.outbox), 0) + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_critical_handler_failures_keep_identity_and_ordinary_errors_translate(self) -> None: + for critical in (KeyboardInterrupt("stop"), SystemExit(7)): + with self.subTest(critical=type(critical).__name__): + runtime = SmallOS().setKernel(Unix()) + + async def critical_handler(socket, error=critical) -> None: + raise error + + state = _make_state( + runtime, _RecordingTransport(), critical_handler + ) + handler_task = SmallTask( + 2, _run_handler, args=(state,), name="critical-handler" + ) + state.handler_task = handler_task + runtime.fork(handler_task) + try: + runtime.start() + except (KeyboardInterrupt, SystemExit) as caught: + self.assertIs(caught, critical) + else: + self.fail("critical handler exception did not escape unchanged") + self.assertIs(state.fatal_error, critical) + self.assertEqual(state.handshake_state, "failed") + + runtime = SmallOS().setKernel(Unix()) + + async def ordinary_handler(socket) -> None: + raise RuntimeError("private detail") + + transport = _RecordingTransport() + state = _make_state(runtime, transport, ordinary_handler) + handler_task = SmallTask( + 2, _run_handler, args=(state,), name="ordinary-handler" + ) + state.handler_task = handler_task + runtime.fork(handler_task) + runtime.start() + self.assertIsNone(handler_task.exception) + self.assertTrue(state.rejected) + self.assertIn(b"HTTP/1.1 500 Internal Server Error", transport.payloads[0]) + self.assertNotIn(b"private detail", transport.payloads[0]) + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_frame_guard_bounds_declared_length_before_payload(self) -> None: + guard = _FrameGuard(max_payload_bytes=1024) + declared = b"\x82\xff" + (65537).to_bytes(8, "big") + b"mask" + with self.assertRaises(WebSocketCapacityError): + for byte in declared: + guard.feed(bytes([byte])) + self.assertLessEqual(len(guard._header), 14) + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_partial_masked_fragmented_input_and_protocol_errors(self) -> None: + api = _load_wsproto() + client = api.Connection(api.ConnectionType.CLIENT) + server = api.Connection(api.ConnectionType.SERVER) + guard = _FrameGuard(1024) + payload = client.send( + api.TextMessage(data="hel", frame_finished=True, message_finished=False) + ) + client.send( + api.TextMessage(data="lo", frame_finished=True, message_finished=True) + ) + for byte in payload: + for chunk in guard.feed(bytes([byte])): + server.receive_data(chunk) + events = list(server.events()) + self.assertEqual("".join(event.data for event in events), "hello") + self.assertTrue(events[-1].message_finished) + with self.assertRaisesRegex(ValueError, "masked"): + _FrameGuard(1024).feed(b"\x81\x01x") + + @unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") + def test_inbound_and_outbound_mailboxes_are_bounded(self) -> None: + config = WebSocketConfig( + max_frame_payload_bytes=8, + max_message_bytes=8, + max_inbound_messages=1, + max_inbound_bytes=3, + max_outbound_commands=1, + max_outbound_bytes=3, + ) + state = _WebSocketState( + object(), + object(), + object(), + upgrade_request(), + _WebSocketRoute(unused_handler, None, ()), + config, + b"", + ) + self.assertTrue(state._deliver_message("abc")) + self.assertFalse(state._deliver_message("x")) + self.assertTrue(state._enqueue_control(state.api.Ping(payload=b"abc"), 3)) + self.assertFalse(state._enqueue_control(state.api.Ping(payload=b"x"), 1)) + + +@unittest.skipUnless(HAS_WSPROTO, "websocket extra is not installed") +class WebSocketLoopbackTests(unittest.TestCase): + def _http_exchange(self, port: int, payload: bytes) -> bytes: + with socket.create_connection(("127.0.0.1", port), timeout=3) as stream: + stream.sendall(payload) + chunks = [] + while True: + chunk = stream.recv(4096) + if not chunk: + return b"".join(chunks) + chunks.append(chunk) + + def test_wsproto_client_interoperability_and_http_coexistence(self) -> None: + api = _load_wsproto() + runtime = SmallOS().setKernel(Unix()) + app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=4096, + max_message_bytes=4096, + idle_timeout=5, + handshake_timeout=2, + close_timeout=1, + ) + ) + + @app.get("/ws") + async def normal_get(request): + return Response.text("ordinary-http") + + @app.websocket( + "/ws", + origins={"https://allowed.example"}, + subprotocols=("chat.v1",), + ) + async def echo(websocket: WebSocket) -> None: + await websocket.accept(subprotocol="chat.v1") + async for message in websocket: + if message.is_text: + await websocket.send_text(message.text) + else: + await websocket.send_bytes(message.bytes) + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + outcomes: list[object] = [] + errors: list[BaseException] = [] + + def client_work() -> None: + try: + ordinary = self._http_exchange( + server.port, + b"GET /ws HTTP/1.1\r\nHost: localhost\r\n\r\n", + ) + outcomes.append(ordinary) + with socket.create_connection( + ("127.0.0.1", server.port), timeout=3 + ) as stream: + stream.sendall( + b"GET /ws?room=1 HTTP/1.1\r\n" + b"Host: localhost\r\n" + b"Upgrade: websocket\r\n" + b"Connection: keep-alive, Upgrade\r\n" + b"Sec-WebSocket-Version: 13\r\n" + b"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n" + b"Sec-WebSocket-Protocol: chat.v1\r\n" + b"Origin: https://allowed.example\r\n\r\n" + ) + response = b"" + while b"\r\n\r\n" not in response: + response += stream.recv(4096) + outcomes.append(response) + client = api.Connection(api.ConnectionType.CLIENT) + stream.sendall( + client.send( + api.TextMessage( + data="hel", + frame_finished=True, + message_finished=False, + ) + ) + + client.send( + api.TextMessage( + data="lo", + frame_finished=True, + message_finished=True, + ) + ) + ) + outcomes.extend(_receive_events(stream, client, api.TextMessage)) + stream.sendall(client.send(api.BytesMessage(data=b"binary"))) + outcomes.extend(_receive_events(stream, client, api.BytesMessage)) + stream.sendall(client.send(api.Ping(payload=b"probe"))) + outcomes.extend(_receive_events(stream, client, api.Pong)) + stream.sendall( + client.send(api.CloseConnection(code=1000, reason="done")) + ) + outcomes.extend(_receive_events(stream, client, api.CloseConnection)) + except BaseException as exc: + errors.append(exc) + finally: + server.close() + + worker = threading.Thread(target=client_work, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=5) + + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertIn(b"ordinary-http", outcomes[0]) + self.assertIn(b"HTTP/1.1 101 Switching Protocols", outcomes[1]) + self.assertIn(b"Sec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo=", outcomes[1]) + self.assertIn(b"Sec-WebSocket-Protocol: chat.v1", outcomes[1]) + self.assertTrue(any(getattr(event, "data", None) == "hello" for event in outcomes)) + self.assertTrue(any(getattr(event, "data", None) == b"binary" for event in outcomes)) + self.assertTrue(any(getattr(event, "payload", None) == b"probe" for event in outcomes)) + self.assertTrue( + any(getattr(event, "code", None) == 1000 for event in outcomes), + outcomes, + ) + self.assertEqual(server._websocket_states, {}) + self.assertEqual(runtime.ioReadWaiters, {}) + self.assertEqual(runtime.ioWriteWaiters, {}) + + def test_smallos_websocket_client_interoperability(self) -> None: + runtime = SmallOS().setKernel(Unix()) + app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=4096, + max_message_bytes=4096, + idle_timeout=5, + close_timeout=1, + ) + ) + + @app.websocket("/native", subprotocols=("smallos.v1",)) + async def echo(websocket: WebSocket) -> None: + await websocket.accept(subprotocol="smallos.v1") + message = await websocket.receive() + await websocket.send_text(message.text) + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + outcome: dict[str, object] = {} + + async def client_job(task) -> None: + client = SmallWebSocketClient( + task, + host="127.0.0.1", + port=server.port, + client_key="dGhlIHNhbXBsZSBub25jZQ==", + ) + try: + await client.connect("/native", subprotocols=("smallos.v1",)) + outcome["subprotocol"] = client.negotiated_subprotocol + await client.send_text("native-client") + outcome["message"] = await client.receive() + finally: + await client.disconnect() + server.close() + + client_task = SmallTask(2, client_job, name="smallserver-ws-client") + runtime.fork(client_task) + runtime.start() + + self.assertIsNone(client_task.exception) + self.assertEqual(outcome["subprotocol"], "smallos.v1") + self.assertEqual( + outcome["message"], {"type": "text", "data": "native-client"} + ) + self.assertTrue(server.finished) + self.assertEqual(runtime.ioReadWaiters, {}) + self.assertEqual(runtime.ioWriteWaiters, {}) + + def test_waiting_websocket_does_not_delay_unrelated_http(self) -> None: + runtime = SmallOS().setKernel(Unix()) + app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=4096, + max_message_bytes=4096, + idle_timeout=5, + close_timeout=1, + ) + ) + + @app.get("/fast") + async def fast(request): + return Response.text("fast") + + @app.websocket("/waiting") + async def waiting(websocket: WebSocket) -> None: + await websocket.accept() + await websocket.receive() + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + api = _load_wsproto() + outcomes: list[bytes] = [] + errors: list[BaseException] = [] + + def client_work() -> None: + try: + with socket.create_connection( + ("127.0.0.1", server.port), timeout=3 + ) as websocket_stream: + websocket_stream.sendall( + b"GET /waiting HTTP/1.1\r\nHost: localhost\r\n" + b"Upgrade: websocket\r\nConnection: Upgrade\r\n" + b"Sec-WebSocket-Version: 13\r\n" + b"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n" + ) + response = b"" + while b"\r\n\r\n" not in response: + response += websocket_stream.recv(4096) + outcomes.append(response) + outcomes.append( + self._http_exchange( + server.port, + b"GET /fast HTTP/1.1\r\nHost: localhost\r\n\r\n", + ) + ) + client = api.Connection(api.ConnectionType.CLIENT) + websocket_stream.sendall( + client.send(api.CloseConnection(code=1000, reason="done")) + ) + _receive_events( + websocket_stream, client, api.CloseConnection + ) + except BaseException as exc: + errors.append(exc) + finally: + server.close() + + worker = threading.Thread(target=client_work, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=5) + + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertIn(b"HTTP/1.1 101 Switching Protocols", outcomes[0]) + self.assertIn(b"fast", outcomes[1]) + self.assertTrue(server.finished) + + def test_malformed_and_oversized_frames_close_with_safe_codes(self) -> None: + runtime = SmallOS().setKernel(Unix()) + app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=8, + max_message_bytes=8, + idle_timeout=5, + close_timeout=1, + ) + ) + + @app.websocket("/bounded") + async def bounded(websocket: WebSocket) -> None: + await websocket.accept() + await websocket.receive() + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + api = _load_wsproto() + close_codes: list[int | None] = [] + errors: list[BaseException] = [] + + def send_bad_frame(frame: bytes) -> None: + with socket.create_connection( + ("127.0.0.1", server.port), timeout=3 + ) as stream: + stream.sendall( + b"GET /bounded HTTP/1.1\r\nHost: localhost\r\n" + b"Upgrade: websocket\r\nConnection: Upgrade\r\n" + b"Sec-WebSocket-Version: 13\r\n" + b"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n" + ) + response = b"" + while b"\r\n\r\n" not in response: + response += stream.recv(4096) + stream.sendall(frame) + client = api.Connection(api.ConnectionType.CLIENT) + events = _receive_events(stream, client, api.CloseConnection) + close = next( + event for event in events if isinstance(event, api.CloseConnection) + ) + close_codes.append(close.code) + stream.sendall(client.send(close.response())) + + def client_work() -> None: + try: + send_bad_frame(b"\x81\x01x") + send_bad_frame(b"\x82\xfe\x00\x7emask") + except BaseException as exc: + errors.append(exc) + finally: + server.close() + + worker = threading.Thread(target=client_work, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=5) + + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertEqual(close_codes, [1002, 1009]) + self.assertTrue(server.finished) + + def test_handshake_idle_and_pong_deadlines_are_bounded(self) -> None: + runtime = SmallOS().setKernel(Unix()) + app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=1024, + max_message_bytes=1024, + handshake_timeout=0.1, + idle_timeout=0.2, + pong_timeout=0.05, + close_timeout=0.2, + deadline_resolution=0.01, + ) + ) + + @app.websocket("/handshake-timeout") + async def handshake_timeout(websocket: WebSocket) -> None: + await runtime.cursor.sleep(1) + + @app.websocket("/idle-timeout") + async def idle_timeout(websocket: WebSocket) -> None: + await websocket.accept() + await runtime.cursor.sleep(5) + + @app.websocket("/pong-timeout") + async def pong_timeout(websocket: WebSocket) -> None: + await websocket.accept() + await websocket.ping(b"deadline") + await websocket.receive() + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + api = _load_wsproto() + outcomes: dict[str, object] = {} + errors: list[BaseException] = [] + + def connect(path: str): + stream = socket.create_connection(("127.0.0.1", server.port), timeout=3) + stream.sendall( + "GET {} HTTP/1.1\r\nHost: localhost\r\n" + "Upgrade: websocket\r\nConnection: Upgrade\r\n" + "Sec-WebSocket-Version: 13\r\n" + "Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n".format( + path + ).encode("ascii") + ) + response = b"" + while b"\r\n\r\n" not in response: + response += stream.recv(4096) + return stream, response + + def client_work() -> None: + try: + stream, response = connect("/handshake-timeout") + outcomes["handshake"] = response + stream.close() + + stream, response = connect("/idle-timeout") + outcomes["idle_handshake"] = response + idle_client = api.Connection(api.ConnectionType.CLIENT) + idle_events = _receive_events( + stream, idle_client, api.CloseConnection + ) + idle_close = next( + event + for event in idle_events + if isinstance(event, api.CloseConnection) + ) + outcomes["idle_code"] = idle_close.code + stream.sendall(idle_client.send(idle_close.response())) + stream.close() + + stream, response = connect("/pong-timeout") + outcomes["pong_handshake"] = response + pong_client = api.Connection(api.ConnectionType.CLIENT) + ping_events = _receive_events(stream, pong_client, api.Ping) + outcomes["ping_payload"] = next( + event.payload + for event in ping_events + if isinstance(event, api.Ping) + ) + close_events = _receive_events( + stream, pong_client, api.CloseConnection + ) + pong_close = next( + event + for event in close_events + if isinstance(event, api.CloseConnection) + ) + outcomes["pong_code"] = pong_close.code + stream.sendall(pong_client.send(pong_close.response())) + stream.close() + except BaseException as exc: + errors.append(exc) + finally: + server.close() + + worker = threading.Thread(target=client_work, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=5) + + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertIn(b"HTTP/1.1 408 Request Timeout", outcomes["handshake"]) + self.assertIn( + b"HTTP/1.1 101 Switching Protocols", outcomes["idle_handshake"] + ) + self.assertEqual(outcomes["idle_code"], 1001) + self.assertEqual(outcomes["ping_payload"], b"deadline") + self.assertEqual(outcomes["pong_code"], 1002) + self.assertTrue(server.finished) + + def test_idle_deadline_cancels_adapter_waiting_handler(self) -> None: + runtime = SmallOS().setKernel(Unix()) + release = threading.Event() + entered = threading.Event() + app = SmallServer( + websocket_config=WebSocketConfig( + idle_timeout=0.05, + close_timeout=0.1, + deadline_resolution=0.01, + ) + ) + errors: list[BaseException] = [] + + def blocking_work() -> None: + entered.set() + if not release.wait(2): + raise TimeoutError("adapter worker was not released") + + with AdapterRegistry( + blocking=ThreadAdapter(max_workers=1, max_pending=1) + ) as services: + + @app.websocket("/adapter-idle") + async def adapter_idle(websocket: WebSocket) -> None: + await websocket.accept() + await services.call("blocking", blocking_work) + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + close_codes: list[int | None] = [] + + def client_work() -> None: + try: + with socket.create_connection( + ("127.0.0.1", server.port), timeout=3 + ) as stream: + stream.sendall( + b"GET /adapter-idle HTTP/1.1\r\nHost: localhost\r\n" + b"Upgrade: websocket\r\nConnection: Upgrade\r\n" + b"Sec-WebSocket-Version: 13\r\n" + b"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n" + ) + response = b"" + while b"\r\n\r\n" not in response: + response += stream.recv(4096) + if not entered.wait(1): + raise TimeoutError("adapter handler did not start") + api = _load_wsproto() + client = api.Connection(api.ConnectionType.CLIENT) + events = _receive_events( + stream, client, api.CloseConnection + ) + close = next( + event + for event in events + if isinstance(event, api.CloseConnection) + ) + close_codes.append(close.code) + stream.sendall(client.send(close.response())) + except BaseException as exc: + errors.append(exc) + finally: + release.set() + server.close() + + worker = threading.Thread(target=client_work, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=5) + + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertEqual(close_codes, [1001]) + self.assertTrue(server.finished) + + def test_server_shutdown_attempts_close_and_releases_children(self) -> None: + api = _load_wsproto() + runtime = SmallOS().setKernel(Unix()) + app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=4096, + max_message_bytes=4096, + close_timeout=1, + idle_timeout=5, + ) + ) + + @app.websocket("/live") + async def live(websocket: WebSocket) -> None: + await websocket.accept() + await websocket.receive() + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + close_events: list[object] = [] + states: list[object] = [] + errors: list[BaseException] = [] + + def client_work() -> None: + try: + with socket.create_connection( + ("127.0.0.1", server.port), timeout=3 + ) as stream: + stream.sendall( + b"GET /live HTTP/1.1\r\nHost: localhost\r\n" + b"Upgrade: websocket\r\nConnection: Upgrade\r\n" + b"Sec-WebSocket-Version: 13\r\n" + b"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n" + ) + response = b"" + while b"\r\n\r\n" not in response: + response += stream.recv(4096) + client = api.Connection(api.ConnectionType.CLIENT) + states.extend(server._websocket_states.values()) + server.close() + events = _receive_events(stream, client, api.CloseConnection) + close_events.extend(events) + close_event = next( + event + for event in events + if isinstance(event, api.CloseConnection) + ) + stream.sendall(client.send(close_event.response())) + except BaseException as exc: + errors.append(exc) + + worker = threading.Thread(target=client_work, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=5) + + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + self.assertTrue( + any(getattr(event, "code", None) == 1001 for event in close_events), + ( + close_events, + states[0].disconnect if states else None, + states[0].handler_error if states else None, + ), + ) + self.assertTrue(server.finished) + self.assertEqual(server._websocket_states, {}) + self.assertEqual(runtime.ioReadWaiters, {}) + self.assertEqual(runtime.ioWriteWaiters, {}) + + def test_handler_failure_sends_sanitized_1011_close(self) -> None: + api = _load_wsproto() + runtime = SmallOS().setKernel(Unix()) + app = SmallServer( + websocket_config=WebSocketConfig( + max_frame_payload_bytes=4096, + max_message_bytes=4096, + close_timeout=1, + idle_timeout=5, + ) + ) + + @app.websocket("/fail") + async def fail(websocket: WebSocket) -> None: + await websocket.accept() + raise RuntimeError("sensitive-handler-detail") + + try: + server = app.serve(runtime, host="127.0.0.1", port=0) + except PermissionError: + self.skipTest("the current sandbox does not permit loopback TCP binds") + + close_events: list[object] = [] + errors: list[BaseException] = [] + + def client_work() -> None: + try: + with socket.create_connection( + ("127.0.0.1", server.port), timeout=3 + ) as stream: + stream.sendall( + b"GET /fail HTTP/1.1\r\nHost: localhost\r\n" + b"Upgrade: websocket\r\nConnection: Upgrade\r\n" + b"Sec-WebSocket-Version: 13\r\n" + b"Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n" + ) + response = b"" + while b"\r\n\r\n" not in response: + response += stream.recv(4096) + client = api.Connection(api.ConnectionType.CLIENT) + events = _receive_events(stream, client, api.CloseConnection) + close_events.extend(events) + close_event = next( + event + for event in events + if isinstance(event, api.CloseConnection) + ) + stream.sendall(client.send(close_event.response())) + except BaseException as exc: + errors.append(exc) + finally: + server.close() + + worker = threading.Thread(target=client_work, daemon=True) + worker.start() + runtime.start() + worker.join(timeout=5) + + self.assertFalse(worker.is_alive()) + self.assertEqual(errors, []) + failure = next( + event for event in close_events if isinstance(event, api.CloseConnection) + ) + self.assertEqual(failure.code, 1011) + self.assertNotIn("sensitive", failure.reason) + self.assertTrue(server.finished) + + +def _receive_events(stream, connection, event_type): + deadline = time.monotonic() + 3 + received = [] + while time.monotonic() < deadline: + data = stream.recv(4096) + if not data: + return received + connection.receive_data(data) + events = list(connection.events()) + received.extend(events) + if any(isinstance(event, event_type) for event in events): + return received + raise TimeoutError("expected WebSocket event was not received") + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/typing/regex_routes.py b/tests/typing/regex_routes.py new file mode 100644 index 0000000..7f66494 --- /dev/null +++ b/tests/typing/regex_routes.py @@ -0,0 +1,20 @@ +"""Public typing fixture for mypy/pyright and compile-only release checks.""" + +from smallserver import Request, Response, RouteErrorEvent, SmallServer + + +def observe(event: RouteErrorEvent) -> None: + route_id: str = event.route_id + assert route_id + + +def application() -> SmallServer: + app = SmallServer(route_error_observer=observe) + + @app.get_regex(r"/users/(?P[0-9]+)") + async def user(request: Request) -> Response: + user_id: str = request.path_params["user_id"] + pattern: str | None = request.route_pattern + return Response.json({"user_id": user_id, "pattern": pattern}) + + return app diff --git a/tests/typing/websocket_routes.py b/tests/typing/websocket_routes.py new file mode 100644 index 0000000..910ecd5 --- /dev/null +++ b/tests/typing/websocket_routes.py @@ -0,0 +1,14 @@ +from smallserver import SmallServer, WebSocket, WebSocketMessage + + +app = SmallServer() + + +@app.websocket("/chat", subprotocols=("chat.v1",)) +async def chat(socket: WebSocket) -> None: + await socket.accept(subprotocol="chat.v1") + message: WebSocketMessage = await socket.receive() + if message.is_text: + await socket.send_text(message.text) + else: + await socket.send_bytes(message.bytes)