From ae72b58ef4c33bf123e9d42ba7d2fa8768681cbb Mon Sep 17 00:00:00 2001 From: jw9829 <245427686+jw9829@users.noreply.github.com> Date: Sun, 27 Sep 2026 20:38:01 -0400 Subject: [PATCH] Make request handlers generic in their server type --- .../@tests/test_cases/check_socketserver.py | 28 +++++++++++++++++++ stdlib/http/server.pyi | 18 +++++++----- stdlib/socketserver.pyi | 28 +++++++++++-------- 3 files changed, 55 insertions(+), 19 deletions(-) create mode 100644 stdlib/@tests/test_cases/check_socketserver.py diff --git a/stdlib/@tests/test_cases/check_socketserver.py b/stdlib/@tests/test_cases/check_socketserver.py new file mode 100644 index 000000000000..fddd1f394290 --- /dev/null +++ b/stdlib/@tests/test_cases/check_socketserver.py @@ -0,0 +1,28 @@ +import socketserver +from http.server import BaseHTTPRequestHandler, HTTPServer, SimpleHTTPRequestHandler +from typing_extensions import assert_type + + +class DefaultHandler(BaseHTTPRequestHandler): + def do_GET(self) -> None: + assert_type(self.server, socketserver.BaseServer) + + +class CustomHTTPServer(HTTPServer): ... + + +class CustomHandler(BaseHTTPRequestHandler[CustomHTTPServer]): + def do_GET(self) -> None: + assert_type(self.server, CustomHTTPServer) + + +class CustomSimpleHandler(SimpleHTTPRequestHandler[CustomHTTPServer]): + def do_GET(self) -> None: + assert_type(self.server, CustomHTTPServer) + + +HTTPServer(("localhost", 0), DefaultHandler) +CustomHTTPServer(("localhost", 0), CustomHandler) +CustomHTTPServer(("localhost", 0), CustomSimpleHandler) +CustomHTTPServer(("localhost", 0), DefaultHandler) +HTTPServer(("localhost", 0), CustomHandler) # type: ignore[arg-type] diff --git a/stdlib/http/server.pyi b/stdlib/http/server.pyi index e739c8a55968..13a63002fda3 100644 --- a/stdlib/http/server.pyi +++ b/stdlib/http/server.pyi @@ -8,7 +8,7 @@ from _typeshed import ReadableBuffer, StrOrBytesPath, StrPath, SupportsRead, Sup from collections.abc import Callable, Iterable, Mapping, Sequence from ssl import Purpose, SSLContext from typing import Any, AnyStr, BinaryIO, ClassVar, Protocol, type_check_only -from typing_extensions import Self, deprecated +from typing_extensions import Self, TypeVar, deprecated __all__ = ["HTTPServer", "ThreadingHTTPServer", "BaseHTTPRequestHandler", "SimpleHTTPRequestHandler"] if sys.version_info < (3, 15): @@ -16,6 +16,8 @@ if sys.version_info < (3, 15): if sys.version_info >= (3, 14): __all__ = ["HTTPSServer", "ThreadingHTTPSServer"] +_ServerT = TypeVar("_ServerT", bound=socketserver.BaseServer, default=socketserver.BaseServer) + class HTTPServer(socketserver.TCPServer): server_name: str server_port: int @@ -43,7 +45,9 @@ if sys.version_info >= (3, 14): def __init__( self, server_address: socketserver._AfInetAddress, - RequestHandlerClass: Callable[[Any, _socket._RetAddress, Self], socketserver.BaseRequestHandler], + RequestHandlerClass: Callable[ + [Any, _socket._RetAddress, Self], socketserver.BaseRequestHandler[Self] | socketserver.BaseRequestHandler + ], bind_and_activate: bool = True, *, certfile: StrOrBytesPath, @@ -55,7 +59,7 @@ if sys.version_info >= (3, 14): class ThreadingHTTPSServer(socketserver.ThreadingMixIn, HTTPSServer): ... -class BaseHTTPRequestHandler(socketserver.StreamRequestHandler): +class BaseHTTPRequestHandler(socketserver.StreamRequestHandler[_ServerT]): client_address: tuple[str, int] close_connection: bool requestline: str @@ -92,7 +96,7 @@ class BaseHTTPRequestHandler(socketserver.StreamRequestHandler): def address_string(self) -> str: ... def parse_request(self) -> bool: ... # undocumented -class SimpleHTTPRequestHandler(BaseHTTPRequestHandler): +class SimpleHTTPRequestHandler(BaseHTTPRequestHandler[_ServerT]): extensions_map: dict[str, str] if sys.version_info >= (3, 12): index_pages: ClassVar[tuple[str, ...]] @@ -102,7 +106,7 @@ class SimpleHTTPRequestHandler(BaseHTTPRequestHandler): self, request: socketserver._RequestType, client_address: _socket._RetAddress, - server: socketserver.BaseServer, + server: _ServerT, *, directory: StrPath | None = None, extra_response_headers: Mapping[str, str] | None = None, @@ -112,7 +116,7 @@ class SimpleHTTPRequestHandler(BaseHTTPRequestHandler): self, request: socketserver._RequestType, client_address: _socket._RetAddress, - server: socketserver.BaseServer, + server: _ServerT, *, directory: StrPath | None = None, ) -> None: ... @@ -129,7 +133,7 @@ def executable(path: StrPath) -> bool: ... # undocumented if sys.version_info < (3, 15): @deprecated("Deprecated and unsafe; removed in Python 3.15.") - class CGIHTTPRequestHandler(SimpleHTTPRequestHandler): + class CGIHTTPRequestHandler(SimpleHTTPRequestHandler[_ServerT]): cgi_directories: list[str] have_fork: bool # undocumented def do_POST(self) -> None: ... diff --git a/stdlib/socketserver.pyi b/stdlib/socketserver.pyi index 05e0025d6a15..f6ca2851c6ce 100644 --- a/stdlib/socketserver.pyi +++ b/stdlib/socketserver.pyi @@ -5,8 +5,8 @@ from _typeshed import ReadableBuffer from collections.abc import Callable from io import BufferedIOBase from socket import socket as _socket -from typing import Any, ClassVar, TypeAlias -from typing_extensions import Self +from typing import Any, ClassVar, Generic, TypeAlias +from typing_extensions import Self, TypeVar __all__ = [ "BaseServer", @@ -41,9 +41,11 @@ _AfInet6Address: TypeAlias = tuple[str | bytes | bytearray, int, int, int] # ad class BaseServer: server_address: _Address timeout: float | None - RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler] + RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler[Self] | BaseRequestHandler] def __init__( - self, server_address: _Address, RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler] + self, + server_address: _Address, + RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler[Self] | BaseRequestHandler], ) -> None: ... def handle_request(self) -> None: ... def serve_forever(self, poll_interval: float = 0.5) -> None: ... @@ -64,6 +66,8 @@ class BaseServer: def shutdown_request(self, request: _RequestType) -> None: ... # undocumented def close_request(self, request: _RequestType) -> None: ... # undocumented +_ServerT = TypeVar("_ServerT", bound=BaseServer, default=BaseServer) + class TCPServer(BaseServer): address_family: int socket: _socket @@ -76,7 +80,7 @@ class TCPServer(BaseServer): def __init__( self, server_address: _AfInetAddress | _AfInet6Address, - RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler], + RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler[Self] | BaseRequestHandler], bind_and_activate: bool = True, ) -> None: ... def fileno(self) -> int: ... @@ -93,7 +97,7 @@ if sys.platform != "win32": def __init__( self, server_address: _AfUnixAddress, - RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler], + RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler[Self] | BaseRequestHandler], bind_and_activate: bool = True, ) -> None: ... @@ -102,7 +106,7 @@ if sys.platform != "win32": def __init__( self, server_address: _AfUnixAddress, - RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler], + RequestHandlerClass: Callable[[Any, _RetAddress, Self], BaseRequestHandler[Self] | BaseRequestHandler], bind_and_activate: bool = True, ) -> None: ... @@ -139,7 +143,7 @@ if sys.platform != "win32": class ThreadingUnixStreamServer(ThreadingMixIn, UnixStreamServer): ... class ThreadingUnixDatagramServer(ThreadingMixIn, UnixDatagramServer): ... -class BaseRequestHandler: +class BaseRequestHandler(Generic[_ServerT]): # `request` is technically of type _RequestType, # but there are some concerns that having a union here would cause # too much inconvenience to people using it (see @@ -148,13 +152,13 @@ class BaseRequestHandler: # Note also that _RetAddress is also just an alias for `Any` request: Any client_address: _RetAddress - server: BaseServer - def __init__(self, request: _RequestType, client_address: _RetAddress, server: BaseServer) -> None: ... + server: _ServerT + def __init__(self, request: _RequestType, client_address: _RetAddress, server: _ServerT) -> None: ... def setup(self) -> None: ... def handle(self) -> None: ... def finish(self) -> None: ... -class StreamRequestHandler(BaseRequestHandler): +class StreamRequestHandler(BaseRequestHandler[_ServerT]): rbufsize: ClassVar[int] # undocumented wbufsize: ClassVar[int] # undocumented timeout: ClassVar[float | None] # undocumented @@ -163,7 +167,7 @@ class StreamRequestHandler(BaseRequestHandler): rfile: BufferedIOBase wfile: BufferedIOBase -class DatagramRequestHandler(BaseRequestHandler): +class DatagramRequestHandler(BaseRequestHandler[_ServerT]): packet: bytes # undocumented socket: _socket # undocumented rfile: BufferedIOBase