Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions stdlib/@tests/test_cases/check_socketserver.py
Original file line number Diff line number Diff line change
@@ -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]
18 changes: 11 additions & 7 deletions stdlib/http/server.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,16 @@ 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):
__all__ += ["CGIHTTPRequestHandler"]
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
Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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, ...]]
Expand All @@ -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,
Expand All @@ -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: ...
Expand All @@ -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: ...
Expand Down
28 changes: 16 additions & 12 deletions stdlib/socketserver.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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: ...
Expand All @@ -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
Expand All @@ -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: ...
Expand All @@ -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: ...

Expand All @@ -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: ...

Expand Down Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down
Loading