actualizado 3-sept

This commit is contained in:
2026-09-03 19:40:44 +02:00
parent 89992823b4
commit 5d2912f789
4338 changed files with 350157 additions and 386 deletions
@@ -0,0 +1,24 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
ASGI support for gunicorn.
This module provides native ASGI worker support, using gunicorn's own
HTTP parsing infrastructure adapted for async I/O.
Components:
- AsyncUnreader: Async socket reading with pushback buffer
- ASGIProtocol: asyncio.Protocol implementation for HTTP handling
- WebSocketProtocol: WebSocket protocol handler (RFC 6455)
- LifespanManager: ASGI lifespan protocol support
Usage:
gunicorn -k asgi myapp:app
"""
from gunicorn.asgi.unreader import AsyncUnreader
from gunicorn.asgi.lifespan import LifespanManager
__all__ = ['AsyncUnreader', 'LifespanManager']
+178
View File
@@ -0,0 +1,178 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
ASGI lifespan protocol manager.
Manages startup and shutdown events for ASGI applications,
enabling frameworks like FastAPI to run initialization and
cleanup code.
"""
import asyncio
class LifespanManager:
"""Manages ASGI lifespan events (startup/shutdown).
The lifespan protocol allows ASGI applications to run code at
startup and shutdown. This is essential for applications that
need to initialize database connections, caches, or other
resources.
ASGI lifespan messages:
- Server sends: {"type": "lifespan.startup"}
- App responds: {"type": "lifespan.startup.complete"} or
{"type": "lifespan.startup.failed", "message": "..."}
- Server sends: {"type": "lifespan.shutdown"}
- App responds: {"type": "lifespan.shutdown.complete"}
"""
def __init__(self, app, logger, state=None):
"""Initialize the lifespan manager.
Args:
app: ASGI application callable
logger: Logger instance
state: Shared state dict for the application
"""
self.app = app
self.logger = logger
self.state = state if state is not None else {}
self._startup_complete = asyncio.Event()
self._shutdown_complete = asyncio.Event()
self._startup_failed = False
self._startup_error = None
self._shutdown_error = None
self._receive_queue = asyncio.Queue()
self._task = None
self._app_finished = False
async def startup(self):
"""Run lifespan startup and wait for completion.
Raises:
RuntimeError: If startup fails or app doesn't support lifespan
"""
scope = {
"type": "lifespan",
"asgi": {"version": "3.0", "spec_version": "2.4"},
"state": self.state,
}
# Send startup event
await self._receive_queue.put({"type": "lifespan.startup"})
# Run lifespan in background task
self._task = asyncio.create_task(self._run_lifespan(scope))
# Wait for startup with timeout
try:
await asyncio.wait_for(
self._startup_complete.wait(),
timeout=30.0 # Reasonable startup timeout
)
except asyncio.TimeoutError:
if self._task:
self._task.cancel()
raise RuntimeError("Lifespan startup timed out")
if self._startup_failed:
if self._task:
self._task.cancel()
msg = self._startup_error or "Unknown error"
raise RuntimeError(f"Lifespan startup failed: {msg}")
self.logger.debug("ASGI lifespan startup complete")
async def shutdown(self):
"""Signal shutdown and wait for completion.
This should be called during graceful shutdown.
"""
if self._app_finished:
self.logger.debug("ASGI lifespan already finished")
return
# Send shutdown event
await self._receive_queue.put({"type": "lifespan.shutdown"})
# Wait for shutdown with timeout
try:
await asyncio.wait_for(
self._shutdown_complete.wait(),
timeout=30.0 # Reasonable shutdown timeout
)
except asyncio.TimeoutError:
self.logger.warning("Lifespan shutdown timed out")
if self._shutdown_error:
self.logger.error("Lifespan shutdown error: %s", self._shutdown_error)
# Cancel the task if still running
if self._task and not self._task.done():
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
self.logger.debug("ASGI lifespan shutdown complete")
async def _run_lifespan(self, scope):
"""Run the ASGI lifespan protocol."""
try:
await self.app(scope, self._receive, self._send)
except asyncio.CancelledError:
raise
except Exception as e:
self.logger.debug("Lifespan application raised: %s", e)
# If startup hasn't completed, mark it as failed
if not self._startup_complete.is_set():
self._startup_failed = True
self._startup_error = str(e)
self._startup_complete.set()
# If shutdown hasn't completed, mark error
elif not self._shutdown_complete.is_set():
self._shutdown_error = str(e)
self._shutdown_complete.set()
finally:
self._app_finished = True
# Ensure events are set to unblock waiters
if not self._startup_complete.is_set():
self._startup_failed = True
self._startup_error = "Application exited before startup complete"
self._startup_complete.set()
if not self._shutdown_complete.is_set():
self._shutdown_complete.set()
async def _receive(self):
"""ASGI receive callable for lifespan."""
return await self._receive_queue.get()
async def _send(self, message):
"""ASGI send callable for lifespan."""
msg_type = message["type"]
if msg_type == "lifespan.startup.complete":
self._startup_complete.set()
self.logger.debug("Received lifespan.startup.complete")
elif msg_type == "lifespan.startup.failed":
self._startup_failed = True
self._startup_error = message.get("message", "")
self._startup_complete.set()
self.logger.debug("Received lifespan.startup.failed: %s",
self._startup_error)
elif msg_type == "lifespan.shutdown.complete":
self._shutdown_complete.set()
self.logger.debug("Received lifespan.shutdown.complete")
elif msg_type == "lifespan.shutdown.failed":
self._shutdown_error = message.get("message", "")
self._shutdown_complete.set()
self.logger.debug("Received lifespan.shutdown.failed: %s",
self._shutdown_error)
+991
View File
@@ -0,0 +1,991 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
HTTP parser for ASGI workers.
Provides callback-based parsing using either the fast C parser (gunicorn_h1c)
or the pure Python PythonProtocol fallback.
"""
import socket
import struct
from enum import IntEnum
class ParseError(Exception):
"""Base error raised during HTTP parsing."""
class InvalidProxyLine(ParseError):
"""Invalid PROXY protocol v1 line."""
class InvalidProxyHeader(ParseError):
"""Invalid PROXY protocol v2 header."""
# PROXY protocol v2 constants
PP_V2_SIGNATURE = b"\x0D\x0A\x0D\x0A\x00\x0D\x0A\x51\x55\x49\x54\x0A"
# RFC 9110 section 6.5.1: fields forbidden in trailers because they alter
# routing, framing, or authentication.
RFC9110_6_5_1_FORBIDDEN_TRAILER = frozenset((
b"host",
b"content-length",
b"transfer-encoding",
b"trailer",
b"authorization",
b"te",
))
class PPCommand(IntEnum):
"""PROXY protocol v2 commands."""
LOCAL = 0x0
PROXY = 0x1
class PPFamily(IntEnum):
"""PROXY protocol v2 address families."""
UNSPEC = 0x0
INET = 0x1 # IPv4
INET6 = 0x2 # IPv6
UNIX = 0x3
class PPProtocol(IntEnum):
"""PROXY protocol v2 transport protocols."""
UNSPEC = 0x0
STREAM = 0x1 # TCP
DGRAM = 0x2 # UDP
class LimitRequestLine(ParseError):
"""Request line exceeds configured limit."""
class LimitRequestHeaders(ParseError):
"""Too many headers or header field too large."""
class InvalidRequestLine(ParseError):
"""Invalid request line."""
class InvalidRequestMethod(ParseError):
"""Invalid HTTP method."""
class InvalidHTTPVersion(ParseError):
"""Invalid HTTP version."""
class InvalidHeaderName(ParseError):
"""Invalid header name."""
class InvalidHeader(ParseError):
"""Invalid header value."""
class UnsupportedTransferCoding(ParseError):
"""Unsupported Transfer-Encoding value."""
class InvalidChunkSize(ParseError):
"""Invalid chunk size in chunked transfer encoding."""
class InvalidChunkExtension(ParseError):
"""Invalid chunk extension per RFC 9112."""
# RFC 9110 section 5.3: fields whose grammar admits only a single member, so a
# second occurrence cannot be merged and makes the message ambiguous. Lowercase
# because that is how _finalize_headers() stores names. content-length belongs
# to this class too but is validated separately, alongside its value.
RFC9110_5_3_SINGLETON_FIELDS = frozenset((
b'host',
b'content-type',
))
class PythonProtocol:
"""Callback-based HTTP/1.1 parser (pure Python fallback).
Mirrors H1CProtocol interface for seamless switching between
the C extension and pure Python implementations.
Callbacks:
on_message_begin: () -> None - Called when request starts
on_url: (url: bytes) -> None - Called with request URL/path
on_header: (name: bytes, value: bytes) -> None - Called for each header
on_headers_complete: () -> bool - Called when headers done (return True to skip body)
on_body: (chunk: bytes) -> None - Called with body data chunks
on_message_complete: () -> None - Called when request is complete
"""
__slots__ = (
'_on_message_begin', '_on_url', '_on_header',
'_on_headers_complete', '_on_body', '_on_message_complete',
'_state', '_buffer', '_headers_list',
'method', 'path', 'http_version', 'headers',
'content_length', 'is_chunked', 'should_keep_alive', 'is_complete',
'_body_remaining', '_skip_body',
'_chunk_state', '_chunk_size', '_chunk_remaining',
'_limit_request_line', '_limit_request_fields', '_limit_request_field_size',
'_permit_unconventional_http_method', '_permit_unconventional_http_version',
'_header_count',
'_proxy_protocol', '_proxy_protocol_info', '_proxy_protocol_done',
)
def __init__(
self,
on_message_begin=None,
on_url=None,
on_header=None,
on_headers_complete=None,
on_body=None,
on_message_complete=None,
limit_request_line=8190,
limit_request_fields=100,
limit_request_field_size=8190,
permit_unconventional_http_method=False,
permit_unconventional_http_version=False,
proxy_protocol='off',
):
self._on_message_begin = on_message_begin
self._on_url = on_url
self._on_header = on_header
self._on_headers_complete = on_headers_complete
self._on_body = on_body
self._on_message_complete = on_message_complete
# Store limits
self._limit_request_line = limit_request_line
self._limit_request_fields = limit_request_fields
self._limit_request_field_size = limit_request_field_size
self._permit_unconventional_http_method = permit_unconventional_http_method
self._permit_unconventional_http_version = permit_unconventional_http_version
self._header_count = 0
# Proxy protocol
self._proxy_protocol = proxy_protocol
self._proxy_protocol_info = None
self._proxy_protocol_done = proxy_protocol == 'off'
# Parser state: proxy_protocol, request_line, headers, body, chunked_size, chunked_data, complete
self._state = 'proxy_protocol' if proxy_protocol != 'off' else 'request_line'
self._buffer = bytearray()
self._headers_list = []
# Request info (populated during parsing)
self.method = None
self.path = None
self.http_version = None
self.headers = []
self.content_length = None
self.is_chunked = False
self.should_keep_alive = True
self.is_complete = False
# Body state
self._body_remaining = 0
self._skip_body = False
# Chunked transfer state
self._chunk_state = 'size' # size, data, trailer
self._chunk_size = 0
self._chunk_remaining = 0
def feed(self, data):
"""Process data, fire callbacks synchronously.
Args:
data: bytes or bytearray of incoming data
Raises:
ParseError: If the HTTP request is malformed
"""
self._buffer.extend(data)
while self._buffer:
if self._state == 'proxy_protocol':
if not self._parse_proxy_protocol():
break
elif self._state == 'request_line':
if not self._parse_request_line():
break
elif self._state == 'headers':
if not self._parse_headers():
break
elif self._state == 'body':
if not self._parse_body():
break
elif self._state == 'chunked':
if not self._parse_chunked_body():
break
else:
break
def remaining(self):
"""Bytes fed after the completed message (b'' if none or not complete).
Matches the accessor H1CProtocol gained in 0.6.8, so a caller does not
have to know which parser it holds. Nothing extra is buffered here:
feed() leaves the state loop once the message completes, and both the
content-length and chunked paths delete what they consume.
"""
if not self.is_complete:
return b''
return bytes(self._buffer)
@property
def remaining_truncated(self):
"""Always False: this parser keeps the whole tail, uncapped."""
return False
@property
def proxy_protocol_info(self):
"""Return proxy protocol info if parsed."""
return self._proxy_protocol_info
def reset(self):
"""Reset for next request (keepalive)."""
self._state = 'request_line'
self._buffer.clear()
self._headers_list = []
self.method = None
self.path = None
self.http_version = None
self.headers = []
self.content_length = None
self.is_chunked = False
self.should_keep_alive = True
self.is_complete = False
self._body_remaining = 0
self._skip_body = False
self._chunk_state = 'size'
self._chunk_size = 0
self._chunk_remaining = 0
self._header_count = 0
def finish(self):
"""Mark parsing complete for EOF handling.
Call when no more data will be received. Handles edge cases like
chunked encoding without final trailer CRLF.
"""
if self._state == 'chunked' and self._chunk_state == 'trailer':
# All body data received, just missing final CRLF
self._state = 'complete'
self.is_complete = True
if self._on_message_complete:
self._on_message_complete()
def _parse_proxy_protocol(self):
"""Parse PROXY protocol header if enabled.
Returns True if parsing is complete (or not applicable),
False if more data is needed.
"""
# Need at least 12 bytes to detect v2 signature or check for v1 prefix
if len(self._buffer) < 12:
return False
mode = self._proxy_protocol
# Check for v2 signature first
if mode in ('v2', 'auto') and self._buffer[:12] == PP_V2_SIGNATURE:
return self._parse_proxy_protocol_v2()
# Check for v1 prefix
if mode in ('v1', 'auto') and self._buffer[:6] == b'PROXY ':
return self._parse_proxy_protocol_v1()
# Not proxy protocol - continue with normal parsing
self._proxy_protocol_done = True
self._state = 'request_line'
return True
def _parse_proxy_protocol_v1(self):
"""Parse PROXY protocol v1 (text format).
Format: PROXY <PROTO> <SRC_ADDR> <DST_ADDR> <SRC_PORT> <DST_PORT>\r\n
"""
# Find end of line
idx = self._buffer.find(b'\r\n')
if idx == -1:
# Need more data - v1 header can be up to 107 bytes
if len(self._buffer) > 107:
raise InvalidProxyLine("PROXY v1 header too long")
return False
line = bytes(self._buffer[:idx]).decode('latin-1')
del self._buffer[:idx + 2]
# Parse the line
parts = line.split(' ')
if len(parts) < 2:
raise InvalidProxyLine("Invalid PROXY v1 line")
proto = parts[1].upper()
if proto == 'UNKNOWN':
# Unknown protocol - no address info
self._proxy_protocol_info = {
'proxy_protocol': 'UNKNOWN',
'client_addr': None,
'client_port': None,
'proxy_addr': None,
'proxy_port': None,
}
elif proto in ('TCP4', 'TCP6'):
if len(parts) != 6:
raise InvalidProxyLine("Invalid PROXY v1 line for %s" % proto)
s_addr = parts[2]
d_addr = parts[3]
# Validate addresses with the appropriate family. WSGI does the
# same in gunicorn/http/message.py:_parse_proxy_protocol_v1.
af = socket.AF_INET if proto == 'TCP4' else socket.AF_INET6
try:
socket.inet_pton(af, s_addr)
socket.inet_pton(af, d_addr)
except (OSError, ValueError):
raise InvalidProxyLine("Invalid PROXY v1 %s address" % proto)
try:
s_port = int(parts[4])
d_port = int(parts[5])
except ValueError as e:
raise InvalidProxyLine("Invalid PROXY v1 port: %s" % e)
if not (0 <= s_port <= 65535 and 0 <= d_port <= 65535):
raise InvalidProxyLine("Invalid PROXY v1 port range")
self._proxy_protocol_info = {
'proxy_protocol': proto,
'client_addr': s_addr,
'client_port': s_port,
'proxy_addr': d_addr,
'proxy_port': d_port,
}
else:
raise InvalidProxyLine("Unknown PROXY v1 protocol: %s" % proto)
self._proxy_protocol_done = True
self._state = 'request_line'
return True
def _parse_proxy_protocol_v2(self):
"""Parse PROXY protocol v2 (binary format)."""
# Need at least 16 bytes for header
if len(self._buffer) < 16:
return False
# Parse header
ver_cmd = self._buffer[12]
fam_prot = self._buffer[13]
length = struct.unpack('>H', bytes(self._buffer[14:16]))[0]
# Check version
version = (ver_cmd & 0xF0) >> 4
if version != 2:
raise InvalidProxyHeader("Unsupported PROXY v2 version: %d" % version)
# Check command
command = ver_cmd & 0x0F
if command not in (PPCommand.LOCAL, PPCommand.PROXY):
raise InvalidProxyHeader("Unsupported PROXY v2 command: %d" % command)
# Check if we have the complete header
total_size = 16 + length
if len(self._buffer) < total_size:
return False
# Extract address data
addr_data = bytes(self._buffer[16:total_size])
del self._buffer[:total_size]
# Handle LOCAL command
if command == PPCommand.LOCAL:
self._proxy_protocol_info = {
'proxy_protocol': 'LOCAL',
'client_addr': None,
'client_port': None,
'proxy_addr': None,
'proxy_port': None,
}
self._proxy_protocol_done = True
self._state = 'request_line'
return True
# Parse address family and protocol
family = (fam_prot & 0xF0) >> 4
protocol = fam_prot & 0x0F
# gunicorn is an HTTP server; only TCP (STREAM) makes sense. WSGI
# rejects non-STREAM at gunicorn/http/message.py:_parse_proxy_protocol_v2.
if family in (PPFamily.INET, PPFamily.INET6) and protocol != PPProtocol.STREAM:
raise InvalidProxyHeader(
"PROXY v2: only TCP (STREAM) protocol is supported"
)
if family == PPFamily.INET:
# IPv4
if len(addr_data) < 12:
raise InvalidProxyHeader("Invalid PROXY v2 IPv4 address data")
s_addr = '.'.join(str(b) for b in addr_data[:4])
d_addr = '.'.join(str(b) for b in addr_data[4:8])
s_port = struct.unpack('>H', addr_data[8:10])[0]
d_port = struct.unpack('>H', addr_data[10:12])[0]
proto = 'TCP4'
elif family == PPFamily.INET6:
# IPv6
if len(addr_data) < 36:
raise InvalidProxyHeader("Invalid PROXY v2 IPv6 address data")
# Format IPv6 addresses
s_words = struct.unpack('>8H', addr_data[:16])
d_words = struct.unpack('>8H', addr_data[16:32])
s_addr = ':'.join('%x' % w for w in s_words)
d_addr = ':'.join('%x' % w for w in d_words)
s_port = struct.unpack('>H', addr_data[32:34])[0]
d_port = struct.unpack('>H', addr_data[34:36])[0]
proto = 'TCP6'
elif family == PPFamily.UNSPEC:
# Unspecified address family
self._proxy_protocol_info = {
'proxy_protocol': 'UNSPEC',
'client_addr': None,
'client_port': None,
'proxy_addr': None,
'proxy_port': None,
}
self._proxy_protocol_done = True
self._state = 'request_line'
return True
else:
raise InvalidProxyHeader("Unsupported PROXY v2 address family: %d" % family)
self._proxy_protocol_info = {
'proxy_protocol': proto,
'client_addr': s_addr,
'client_port': s_port,
'proxy_addr': d_addr,
'proxy_port': d_port,
}
self._proxy_protocol_done = True
self._state = 'request_line'
return True
def _parse_request_line(self):
"""Parse request line, return True if complete."""
idx = self._buffer.find(b'\r\n')
if idx == -1:
return False
# Check request line length limit
if self._limit_request_line > 0 and idx > self._limit_request_line:
raise LimitRequestLine("Request line is too large")
line = bytes(self._buffer[:idx])
del self._buffer[:idx + 2]
# Parse: METHOD PATH HTTP/x.y
parts = line.split(b' ', 2)
if len(parts) != 3:
raise InvalidRequestLine("Invalid request line")
self.method = parts[0]
self.path = parts[1]
# Validate method
if not self._permit_unconventional_http_method:
if not self._is_valid_method(self.method):
raise InvalidRequestMethod(self.method.decode('latin-1'))
# RFC 9112 section 3.2.4: asterisk-form is only valid with OPTIONS.
if self.path == b'*' and self.method != b'OPTIONS':
raise InvalidRequestLine("Invalid request line")
# RFC 9112 section 3.2.3: authority-form is only valid with CONNECT.
if (self.method != b'CONNECT'
and self.path != b'*'
and not self.path.startswith(b'/')
and b'://' not in self.path):
raise InvalidRequestLine("Invalid request line")
# Parse version
version = parts[2]
if version == b'HTTP/1.1':
self.http_version = (1, 1)
elif version == b'HTTP/1.0':
self.http_version = (1, 0)
else:
if not self._permit_unconventional_http_version:
raise InvalidHTTPVersion(version.decode('latin-1'))
# Try to parse other HTTP/1.x versions if permitted
if version.startswith(b'HTTP/1.'):
try:
minor = int(version[7:])
self.http_version = (1, minor)
except ValueError:
raise InvalidHTTPVersion(version.decode('latin-1'))
else:
raise InvalidHTTPVersion(version.decode('latin-1'))
if self._on_message_begin:
self._on_message_begin()
if self._on_url:
self._on_url(self.path)
self._state = 'headers'
return True
def _parse_headers(self):
"""Parse headers, return True if headers are complete."""
while True:
idx = self._buffer.find(b'\r\n')
if idx == -1:
return False
line = bytes(self._buffer[:idx])
del self._buffer[:idx + 2]
if not line:
# Empty line = end of headers
self._finalize_headers()
return True
# Check header field size limit (include CRLF in size to match WSGI parser)
if self._limit_request_field_size > 0 and len(line) + 2 > self._limit_request_field_size:
raise LimitRequestHeaders("Request header field is too large")
# Check header count limit
self._header_count += 1
if self._limit_request_fields > 0 and self._header_count > self._limit_request_fields:
raise LimitRequestHeaders("Too many headers")
# Parse header
colon = line.find(b':')
if colon == -1:
raise InvalidHeader("Missing colon in header")
name = line[:colon].strip()
if not self._is_valid_token(name):
raise InvalidHeaderName(name.decode('latin-1'))
value = line[colon + 1:].strip()
if self._has_invalid_header_chars(value):
raise InvalidHeader("Invalid characters in header value")
# Store lowercase name for internal use
name_lower = name.lower()
self._headers_list.append((name_lower, value))
if self._on_header:
self._on_header(name_lower, value)
def _finalize_headers(self):
"""Called when all headers received.
Validates headers for request smuggling vulnerabilities:
- Rejects duplicate Content-Length headers
- Rejects duplicate Host and Content-Type headers
- Rejects requests with both Content-Length and Transfer-Encoding
- Rejects chunked Transfer-Encoding in HTTP/1.0
- Rejects stacked chunked encoding
- Validates Transfer-Encoding values
"""
self.headers = self._headers_list
# Extract and validate content-length and transfer-encoding
content_length = None
chunked = False
seen_singletons = set()
for name, value in self.headers:
# RFC 9110 section 5.3: these admit a single member only, so a
# repeat cannot be merged and leaves the message ambiguous.
# content-length is handled separately just below.
if name in RFC9110_5_3_SINGLETON_FIELDS:
if name in seen_singletons:
raise InvalidHeader(
"Duplicate %s header" % name.decode('latin-1'))
seen_singletons.add(name)
if name == b'content-length':
# Reject duplicate Content-Length headers (request smuggling vector)
if content_length is not None:
raise InvalidHeader("Duplicate Content-Length header")
try:
cl_value = int(value)
except ValueError:
raise InvalidHeader("Invalid Content-Length value")
if cl_value < 0:
raise InvalidHeader("Negative Content-Length")
content_length = cl_value
elif name == b'transfer-encoding':
# Properly parse comma-separated Transfer-Encoding values
# per RFC 9112 Section 6.1
vals = [v.strip() for v in value.split(b',')]
for val in vals:
val_lower = val.lower()
if val_lower == b'chunked':
# Reject stacked chunked encoding (request smuggling vector)
if chunked:
raise InvalidHeader("Stacked chunked encoding")
chunked = True
elif val_lower == b'identity':
# identity after chunked is invalid
if chunked:
raise InvalidHeader("Invalid Transfer-Encoding after chunked")
elif val_lower in (b'compress', b'deflate', b'gzip'):
# Compression after chunked is invalid
if chunked:
raise InvalidHeader("Invalid Transfer-Encoding after chunked")
# Mark connection for close (unsupported but valid)
self.should_keep_alive = False
else:
# Reject unknown transfer codings
raise UnsupportedTransferCoding(val.decode('latin-1'))
elif name == b'connection':
val = value.lower()
if b'close' in val:
self.should_keep_alive = False
elif b'keep-alive' in val:
self.should_keep_alive = True
# Security checks for request smuggling prevention
if chunked:
# Reject chunked in HTTP/1.0 (RFC 9112 Section 6.1)
if self.http_version < (1, 1):
raise InvalidHeader("Chunked encoding not allowed in HTTP/1.0")
# Reject Content-Length with Transfer-Encoding (request smuggling vector)
if content_length is not None:
raise InvalidHeader("Content-Length with Transfer-Encoding")
self.is_chunked = True
self.content_length = None
self._body_remaining = -1 # Chunked mode
elif content_length is not None:
self.content_length = content_length
self._body_remaining = content_length
else:
# No body
self.content_length = None
self._body_remaining = 0
# HTTP/1.0 defaults to close
if self.http_version == (1, 0) and self.should_keep_alive:
# Only keep-alive if explicitly requested
has_keepalive = any(
name == b'connection' and b'keep-alive' in value.lower()
for name, value in self.headers
)
if not has_keepalive:
self.should_keep_alive = False
if self._on_headers_complete:
self._skip_body = self._on_headers_complete()
# Determine next state
if self._skip_body:
self._state = 'complete'
self.is_complete = True
if self._on_message_complete:
self._on_message_complete()
elif self.is_chunked:
self._state = 'chunked'
self._chunk_state = 'size'
elif self.content_length and self.content_length > 0:
self._state = 'body'
else:
# No body
self._state = 'complete'
self.is_complete = True
if self._on_message_complete:
self._on_message_complete()
def _parse_body(self):
"""Parse Content-Length delimited body."""
if not self._buffer or self._body_remaining <= 0:
return False
chunk_size = min(len(self._buffer), self._body_remaining)
chunk = bytes(self._buffer[:chunk_size])
del self._buffer[:chunk_size]
self._body_remaining -= chunk_size
if self._on_body:
self._on_body(chunk)
if self._body_remaining <= 0:
self._state = 'complete'
self.is_complete = True
if self._on_message_complete:
self._on_message_complete()
return True
def _parse_chunked_body(self):
"""Parse chunked transfer encoding."""
while self._buffer:
if self._chunk_state == 'size':
# Looking for chunk size line
idx = self._buffer.find(b'\r\n')
if idx == -1:
return False
size_line = bytes(self._buffer[:idx])
del self._buffer[:idx + 2]
# Handle chunk extensions (e.g., "5;ext=value")
semicolon = size_line.find(b';')
if semicolon != -1:
# RFC 9112: chunk-ext must not contain bare CR
chunk_ext = size_line[semicolon + 1:]
if b'\r' in chunk_ext:
raise InvalidChunkExtension("bare CR not allowed")
size_line = size_line[:semicolon]
# Strict validation: reject leading/trailing whitespace
# to prevent parser desync (request smuggling vector)
if size_line != size_line.strip():
raise InvalidChunkSize("Whitespace in chunk size")
if not size_line:
raise InvalidChunkSize("Empty chunk size")
# Validate hex characters only (0-9, a-f, A-F)
for c in size_line:
if c not in b'0123456789abcdefABCDEF':
raise InvalidChunkSize("Invalid character in chunk size")
try:
self._chunk_size = int(size_line, 16)
except ValueError:
raise InvalidChunkSize("Invalid chunk size")
if self._chunk_size == 0:
# Final chunk - skip trailers
self._chunk_state = 'trailer'
else:
self._chunk_remaining = self._chunk_size
self._chunk_state = 'data'
elif self._chunk_state == 'data':
# Reading chunk data
if not self._buffer:
return False
to_read = min(len(self._buffer), self._chunk_remaining)
chunk = bytes(self._buffer[:to_read])
del self._buffer[:to_read]
self._chunk_remaining -= to_read
if self._on_body:
self._on_body(chunk)
if self._chunk_remaining == 0:
# Need to consume trailing CRLF
self._chunk_state = 'crlf'
elif self._chunk_state == 'crlf':
# Skip CRLF after chunk data
if len(self._buffer) < 2:
return False
del self._buffer[:2] # Skip \r\n
self._chunk_state = 'size'
elif self._chunk_state == 'trailer':
# Skip trailer headers
idx = self._buffer.find(b'\r\n')
if idx == -1:
return False
line = bytes(self._buffer[:idx])
del self._buffer[:idx + 2]
if not line:
# Empty line = end of trailers
self._state = 'complete'
self.is_complete = True
if self._on_message_complete:
self._on_message_complete()
return True
# RFC 9110 section 6.5.1: reject fields that must not appear
# in trailers.
colon = line.find(b':')
if colon > 0:
name = line[:colon].strip(b' \t').lower()
if name in RFC9110_6_5_1_FORBIDDEN_TRAILER:
raise InvalidHeaderName(name.decode('latin-1'))
return False
def _is_valid_method(self, method):
"""Check if method is valid token with conventional restrictions."""
if not method:
return False
# Check length (3-20 chars)
if not 3 <= len(method) <= 20:
return False
# Check for lowercase or # (unconventional)
for c in method:
if c in b'abcdefghijklmnopqrstuvwxyz#':
return False
return self._is_valid_token(method)
def _is_valid_token(self, data):
"""Check if data contains only RFC 9110 token characters."""
if not data:
return False
for c in data:
if c < 0x21 or c > 0x7e:
return False
# RFC 9110 delimiters: "(),/:;<=>?@[\]{}
if c in b'"(),/:;<=>?@[\\]{}"':
return False
return True
def _has_invalid_header_chars(self, value):
"""RFC 9110 section 5.5: only VCHAR, SP, HTAB, and obs-text allowed."""
for c in value:
if c <= 0x08 or 0x0a <= c <= 0x1f or c == 0x7f:
return True
return False
class CallbackRequest:
"""Request object built from callback parser state.
Works with both H1CProtocol (C extension) and PythonProtocol.
"""
__slots__ = (
'method', 'uri', 'path', 'query', 'fragment', 'version',
'headers', 'headers_bytes', 'scheme', 'raw_path',
'content_length', 'chunked', 'must_close',
'proxy_protocol_info', '_expect_100_continue',
)
def __init__(self):
self.method = None
self.uri = None
self.path = None
self.query = None
self.fragment = None
self.version = None
self.headers = []
self.headers_bytes = []
self.scheme = "http"
self.raw_path = b''
self.content_length = 0
self.chunked = False
self.must_close = False
self.proxy_protocol_info = None
self._expect_100_continue = False
@classmethod
def from_parser(cls, parser, is_ssl=False):
"""Build request from callback parser state.
Args:
parser: H1CProtocol or PythonProtocol instance
is_ssl: Whether connection is SSL/TLS
Returns:
CallbackRequest instance
"""
from urllib.parse import unquote_to_bytes
req = cls()
req.method = parser.method.decode('ascii')
# Parse path and query from URL
# Per ASGI spec:
# - path: percent-decoded UTF-8 string
# - raw_path: original bytes as received
raw_url = parser.path
if b'?' in raw_url:
path_part, query_part = raw_url.split(b'?', 1)
req.raw_path = path_part # Store original bytes
req.path = unquote_to_bytes(path_part).decode('utf-8', errors='replace')
req.query = query_part.decode('latin-1')
else:
req.raw_path = raw_url # Store original bytes
req.path = unquote_to_bytes(raw_url).decode('utf-8', errors='replace')
req.query = ''
req.uri = raw_url.decode('latin-1')
req.fragment = ''
req.version = parser.http_version
# Headers - store both bytes (for ASGI scope) and strings (for compatibility)
# Use asgi_headers (lowercase names) if available (fast parser >= 0.6.2),
# otherwise fall back to headers (Python parser already uses lowercase)
req.headers_bytes = list(getattr(parser, 'asgi_headers', None) or parser.headers)
req.headers = [
(n.decode('latin-1').upper(), v.decode('latin-1'))
for n, v in parser.headers
]
# RFC 9110 section 5.3, enforced here because this is where both
# parsers converge. Both reject these on their own now (PythonProtocol
# in _finalize_headers(), H1CProtocol since 0.6.6), so this is a
# backstop: the pip requirement is not enforced at runtime, and an
# older gunicorn_h1c would otherwise let duplicates through.
seen_singletons = set()
for name, _value in parser.headers:
lowered = name.lower()
if lowered in RFC9110_5_3_SINGLETON_FIELDS:
if lowered in seen_singletons:
raise InvalidHeader(
"Duplicate %s header" % lowered.decode('latin-1'))
seen_singletons.add(lowered)
req.scheme = 'https' if is_ssl else 'http'
req.content_length = parser.content_length or 0
req.chunked = parser.is_chunked
req.must_close = not parser.should_keep_alive
# Check for Expect: 100-continue
for name, value in parser.headers:
if name == b'expect' and value.lower() == b'100-continue':
req._expect_100_continue = True
break
return req
def should_close(self):
"""Check if connection should be closed after this request."""
if self.must_close:
return True
for name, value in self.headers:
if name == "CONNECTION":
v = value.lower().strip(" \t")
if v == "close":
return True
elif v == "keep-alive":
return False
break
return self.version <= (1, 0)
def get_header(self, name):
"""Get a header value by name (case-insensitive)."""
name = name.upper()
for h, v in self.headers:
if h == name:
return v
return None
File diff suppressed because it is too large Load Diff
+135
View File
@@ -0,0 +1,135 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Async version of gunicorn/http/unreader.py for ASGI workers.
Provides async reading with pushback buffer support.
"""
import io
class AsyncUnreader:
"""Async socket reader with pushback buffer support.
This class wraps an asyncio StreamReader and provides the ability
to "unread" data back into a buffer for re-parsing.
Performance optimization: Reuses BytesIO buffer with truncate/seek
instead of creating new objects to reduce GC pressure.
"""
def __init__(self, reader, max_chunk=8192):
"""Initialize the async unreader.
Args:
reader: asyncio.StreamReader instance
max_chunk: Maximum bytes to read at once
"""
self.reader = reader
self.buf = io.BytesIO()
self.max_chunk = max_chunk
self._buf_start = 0 # Start position of valid data in buffer
def _reset_buffer(self):
"""Reset buffer for reuse instead of creating new BytesIO."""
self.buf.seek(0)
self.buf.truncate(0)
self._buf_start = 0
def _get_buffered_data(self):
"""Get all buffered data and reset buffer."""
self.buf.seek(self._buf_start)
data = self.buf.read()
self._reset_buffer()
return data
def _buffer_size(self):
"""Get size of buffered data."""
end = self.buf.seek(0, io.SEEK_END)
return end - self._buf_start
async def read(self, size=None):
"""Read data from the stream, using buffered data first.
Args:
size: Number of bytes to read. If None, returns all buffered
data or reads a single chunk.
Returns:
bytes: Data read from buffer or stream
"""
if size is not None and not isinstance(size, int):
raise TypeError("size parameter must be an int or long.")
if size is not None:
if size == 0:
return b""
if size < 0:
size = None
buf_size = self._buffer_size()
# If no size specified, return buffered data or read chunk
if size is None and buf_size > 0:
return self._get_buffered_data()
if size is None:
chunk = await self._read_chunk()
return chunk
# Read until we have enough data
while buf_size < size:
chunk = await self._read_chunk()
if not chunk:
return self._get_buffered_data()
self.buf.seek(0, io.SEEK_END)
self.buf.write(chunk)
buf_size += len(chunk)
# We have enough data - extract what we need
self.buf.seek(self._buf_start)
data = self.buf.read(size)
# Update start position instead of creating new buffer
self._buf_start += size
# If buffer is getting large with consumed data, compact it
if self._buf_start > 8192:
remaining = self.buf.read() # Read from current position
self._reset_buffer()
if remaining:
self.buf.write(remaining)
return data
async def _read_chunk(self):
"""Read a chunk of data from the underlying stream."""
try:
return await self.reader.read(self.max_chunk)
except Exception:
return b""
def unread(self, data):
"""Push data back into the buffer for re-reading.
Args:
data: bytes to push back
Note: This prepends data to the buffer so it will be read first.
"""
if data:
# Get existing buffered data
self.buf.seek(self._buf_start)
existing = self.buf.read()
# Reset and write new data first, then existing
self._reset_buffer()
self.buf.write(data)
if existing:
self.buf.write(existing)
def has_buffered_data(self):
"""Check if there's data in the pushback buffer."""
return self._buffer_size() > 0
+172
View File
@@ -0,0 +1,172 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""Async uWSGI protocol parser for ASGI workers.
Reuses the parsing logic from gunicorn/uwsgi/message.py, only async I/O differs.
"""
from gunicorn.uwsgi.message import UWSGIRequest
from gunicorn.uwsgi.errors import (
InvalidUWSGIHeader,
UnsupportedModifier,
)
class AsyncUWSGIRequest(UWSGIRequest):
"""Async version of UWSGIRequest.
Reuses all parsing logic from the sync version, only async I/O differs.
The following methods are reused from the parent class:
- _parse_vars() - pure parsing, no I/O
- _extract_request_info() - pure transformation
- _check_allowed_ip() - no I/O
- should_close() - simple logic
"""
# pylint: disable=super-init-not-called
def __init__(self, cfg, unreader, peer_addr, req_number=1):
# Don't call super().__init__ - it does sync parsing
# Just initialize attributes
self.cfg = cfg
self.unreader = unreader
self.peer_addr = peer_addr
self.remote_addr = peer_addr
self.req_number = req_number
# Initialize all attributes (same as sync version)
self.method = None
self.uri = None
self.path = None
self.query = None
self.fragment = ""
self.version = (1, 1)
self.headers = []
self.trailers = []
self.body = None
self.scheme = "https" if cfg.is_ssl else "http"
self.must_close = False
self.uwsgi_vars = {}
self.modifier1 = 0
self.modifier2 = 0
self.proxy_protocol_info = None
# Body state
self.content_length = 0
self.chunked = False
self._body_remaining = 0
# Async factory method - intentionally differs from sync parent:
# - async instead of sync (invalid-overridden-method)
# - different signature for async I/O (arguments-differ)
# pylint: disable=arguments-differ,invalid-overridden-method
@classmethod
async def parse(cls, cfg, unreader, peer_addr, req_number=1):
"""Parse a uWSGI request asynchronously.
Args:
cfg: gunicorn config object
unreader: AsyncUnreader instance
peer_addr: client address tuple
req_number: request number on this connection (for keepalive)
Returns:
AsyncUWSGIRequest: Parsed request object
Raises:
InvalidUWSGIHeader: If the uWSGI header is malformed
UnsupportedModifier: If modifier1 is not 0
ForbiddenUWSGIRequest: If source IP is not allowed
"""
req = cls(cfg, unreader, peer_addr, req_number)
req._check_allowed_ip() # Reuse from parent
await req._async_parse()
return req
async def _async_parse(self):
"""Async version of parse() - reads data then uses sync parsing."""
# Read 4-byte header
header = await self._async_read_exact(4)
if len(header) < 4:
raise InvalidUWSGIHeader("incomplete header")
self.modifier1 = header[0]
datasize = int.from_bytes(header[1:3], 'little')
self.modifier2 = header[3]
if self.modifier1 != 0:
raise UnsupportedModifier(self.modifier1)
# Read vars block
if datasize > 0:
vars_data = await self._async_read_exact(datasize)
if len(vars_data) < datasize:
raise InvalidUWSGIHeader("incomplete vars block")
self._parse_vars(vars_data) # Reuse sync method
self._extract_request_info() # Reuse sync method
self._set_body_reader()
async def _async_read_exact(self, size):
"""Read exactly size bytes asynchronously."""
buf = bytearray()
while len(buf) < size:
chunk = await self.unreader.read(size - len(buf))
if not chunk:
break
buf.extend(chunk)
return bytes(buf)
def _set_body_reader(self):
"""Set up body state for async reading."""
content_length = 0
if 'CONTENT_LENGTH' in self.uwsgi_vars:
try:
content_length = max(int(self.uwsgi_vars['CONTENT_LENGTH']), 0)
except ValueError:
content_length = 0
self.content_length = content_length
self._body_remaining = content_length
async def read_body(self, size=8192):
"""Read body chunk asynchronously.
Args:
size: Maximum bytes to read
Returns:
bytes: Body data, empty bytes when body is exhausted
"""
if self._body_remaining <= 0:
return b""
to_read = min(size, self._body_remaining)
data = await self.unreader.read(to_read)
if data:
self._body_remaining -= len(data)
return data
async def drain_body(self):
"""Drain unread body data.
Should be called before reusing connection for keepalive.
"""
while self._body_remaining > 0:
data = await self.read_body(8192)
if not data:
break
def get_header(self, name):
"""Get header by name (case-insensitive).
Args:
name: Header name to look up
Returns:
Header value if found, None otherwise
"""
name = name.upper()
for h, v in self.headers:
if h == name:
return v
return None
@@ -0,0 +1,437 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
WebSocket protocol handler for ASGI.
Implements RFC 6455 WebSocket protocol for ASGI applications.
"""
import asyncio
import base64
import hashlib
import struct
# WebSocket frame opcodes
OPCODE_CONTINUATION = 0x0
OPCODE_TEXT = 0x1
OPCODE_BINARY = 0x2
OPCODE_CLOSE = 0x8
OPCODE_PING = 0x9
OPCODE_PONG = 0xA
# WebSocket close codes
CLOSE_NORMAL = 1000
CLOSE_GOING_AWAY = 1001
CLOSE_PROTOCOL_ERROR = 1002
CLOSE_UNSUPPORTED = 1003
CLOSE_NO_STATUS = 1005
CLOSE_ABNORMAL = 1006
CLOSE_INVALID_DATA = 1007
CLOSE_POLICY_VIOLATION = 1008
CLOSE_MESSAGE_TOO_BIG = 1009
CLOSE_MANDATORY_EXT = 1010
CLOSE_INTERNAL_ERROR = 1011
# WebSocket handshake GUID (RFC 6455)
WS_GUID = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
class WebSocketProtocol:
"""WebSocket connection handler for ASGI applications.
Uses callback-based data feeding instead of StreamReader for efficiency.
Data is fed via feed_data() from the parent protocol's data_received().
"""
def __init__(self, transport, scope, app, log):
"""Initialize WebSocket protocol handler.
Args:
transport: asyncio transport for writing
scope: ASGI WebSocket scope dict
app: ASGI application callable
log: Logger instance
"""
self.transport = transport
self.scope = scope
self.app = app
self.log = log
self.accepted = False
self.closed = False
self.close_code = None
self.close_reason = ""
# Close handshake state (RFC 6455 Section 7.1.1)
self._close_sent = False
self._close_received = False
self._close_event = asyncio.Event()
# Message reassembly state
self._fragments = []
self._fragment_opcode = None
# Receive queue for incoming messages
self._receive_queue = asyncio.Queue()
# Callback-based data reception (replaces StreamReader)
self._buffer = bytearray()
self._data_event = asyncio.Event()
self._eof = False
def feed_data(self, data):
"""Feed incoming data from the parent protocol's data_received().
Args:
data: bytes received on the connection
"""
if data:
self._buffer.extend(data)
self._data_event.set()
def feed_eof(self):
"""Signal that the connection has been closed."""
self._eof = True
self._data_event.set()
async def run(self):
"""Run the WebSocket ASGI application."""
# Send initial connect event
await self._receive_queue.put({"type": "websocket.connect"})
# Start frame reading task
read_task = asyncio.create_task(self._read_frames())
try:
await self.app(self.scope, self._receive, self._send)
except Exception:
self.log.exception("Error in WebSocket ASGI application")
finally:
# Send close frame if not already closed
if not self.closed and self.accepted and not self._close_sent:
await self._send_close(CLOSE_INTERNAL_ERROR, "Application error")
# Wait for client's close response
try:
await asyncio.wait_for(self._close_event.wait(), timeout=5.0)
except asyncio.TimeoutError:
self.closed = True
read_task.cancel()
try:
await read_task
except asyncio.CancelledError:
pass
async def _receive(self):
"""ASGI receive callable."""
return await self._receive_queue.get()
async def _send(self, message):
"""ASGI send callable."""
msg_type = message["type"]
if msg_type == "websocket.accept":
if self.accepted:
raise RuntimeError("WebSocket already accepted")
await self._send_accept(message)
self.accepted = True
elif msg_type == "websocket.send":
if not self.accepted:
raise RuntimeError("WebSocket not accepted")
if self.closed:
raise RuntimeError("WebSocket closed")
# Check for truthy values since both keys may be present with None
text = message.get("text")
bytes_data = message.get("bytes")
if text is not None:
await self._send_frame(OPCODE_TEXT, text.encode("utf-8"))
elif bytes_data is not None:
await self._send_frame(OPCODE_BINARY, bytes_data)
elif msg_type == "websocket.close":
code = message.get("code", CLOSE_NORMAL)
reason = message.get("reason", "")
await self._send_close(code, reason)
# Wait for client's close frame (RFC 6455 close handshake)
try:
await asyncio.wait_for(self._close_event.wait(), timeout=5.0)
except asyncio.TimeoutError:
self.log.debug("WebSocket close handshake timeout")
self.closed = True
self._close_event.set()
# Close the transport after close handshake
self.transport.close()
async def _send_accept(self, message):
"""Send WebSocket handshake accept response."""
# Get Sec-WebSocket-Key from headers
ws_key = None
for name, value in self.scope["headers"]:
if name == b"sec-websocket-key":
ws_key = value
break
if not ws_key:
raise RuntimeError("Missing Sec-WebSocket-Key header")
# Calculate accept key
accept_key = base64.b64encode(
hashlib.sha1(ws_key + WS_GUID).digest()
).decode("ascii")
# Build response headers
headers = [
"HTTP/1.1 101 Switching Protocols\r\n",
"Upgrade: websocket\r\n",
"Connection: Upgrade\r\n",
f"Sec-WebSocket-Accept: {accept_key}\r\n",
]
# Add selected subprotocol if specified
subprotocol = message.get("subprotocol")
if subprotocol:
headers.append(f"Sec-WebSocket-Protocol: {subprotocol}\r\n")
# Add any extra headers from message
extra_headers = message.get("headers", [])
for name, value in extra_headers:
if isinstance(name, bytes):
name = name.decode("latin-1")
if isinstance(value, bytes):
value = value.decode("latin-1")
headers.append(f"{name}: {value}\r\n")
headers.append("\r\n")
self.transport.write("".join(headers).encode("latin-1"))
async def _read_frames(self):
"""Read and process incoming WebSocket frames."""
try:
# Continue reading while not closed, or if we sent close but haven't
# received client's close response yet (RFC 6455 close handshake)
while not self.closed or (self._close_sent and not self._close_received):
frame = await self._read_frame()
if frame is None:
break
opcode, payload = frame
if opcode == OPCODE_CLOSE:
await self._handle_close(payload)
break
if opcode == OPCODE_PING:
await self._send_frame(OPCODE_PONG, payload)
elif opcode == OPCODE_PONG:
# Ignore pongs
pass
elif opcode == OPCODE_TEXT:
await self._receive_queue.put({
"type": "websocket.receive",
"text": payload.decode("utf-8"),
})
elif opcode == OPCODE_BINARY:
await self._receive_queue.put({
"type": "websocket.receive",
"bytes": payload,
})
elif opcode == OPCODE_CONTINUATION:
# Handle fragmented messages
await self._handle_continuation(payload)
except asyncio.CancelledError:
raise
except Exception as e:
self.log.debug("WebSocket read error: %s", e)
finally:
# Signal disconnect
if not self.closed:
self.closed = True
await self._receive_queue.put({
"type": "websocket.disconnect",
"code": self.close_code or CLOSE_ABNORMAL,
})
async def _read_frame(self): # pylint: disable=too-many-return-statements
"""Read a single WebSocket frame.
Returns:
tuple: (opcode, payload) or None if connection closed
"""
# Read frame header (2 bytes minimum)
header = await self._read_exact(2)
if not header:
return None
first_byte, second_byte = header[0], header[1]
fin = (first_byte >> 7) & 1
rsv1 = (first_byte >> 6) & 1
rsv2 = (first_byte >> 5) & 1
rsv3 = (first_byte >> 4) & 1
opcode = first_byte & 0x0F
# RSV bits must be 0 (no extensions)
if rsv1 or rsv2 or rsv3:
await self._send_close(CLOSE_PROTOCOL_ERROR, "RSV bits set")
return None
masked = (second_byte >> 7) & 1
payload_len = second_byte & 0x7F
# Client frames must be masked (RFC 6455)
if not masked:
await self._send_close(CLOSE_PROTOCOL_ERROR, "Frame not masked")
return None
# Extended payload length
if payload_len == 126:
ext_len = await self._read_exact(2)
if not ext_len:
return None
payload_len = struct.unpack("!H", ext_len)[0]
elif payload_len == 127:
ext_len = await self._read_exact(8)
if not ext_len:
return None
payload_len = struct.unpack("!Q", ext_len)[0]
# Read masking key
masking_key = await self._read_exact(4)
if not masking_key:
return None
# Read payload
payload = await self._read_exact(payload_len)
if payload is None:
return None
# Unmask payload
payload = self._unmask(payload, masking_key)
# Handle fragmented messages
if opcode == OPCODE_CONTINUATION:
if self._fragment_opcode is None:
await self._send_close(CLOSE_PROTOCOL_ERROR, "Unexpected continuation")
return None
self._fragments.append(payload)
if fin:
# Reassemble complete message
full_payload = b"".join(self._fragments)
final_opcode = self._fragment_opcode
self._fragments = []
self._fragment_opcode = None
return (final_opcode, full_payload)
return (OPCODE_CONTINUATION, b"") # Fragment received, wait for more
elif opcode in (OPCODE_TEXT, OPCODE_BINARY):
if not fin:
# Start of fragmented message
self._fragment_opcode = opcode
self._fragments = [payload]
return (OPCODE_CONTINUATION, b"") # Fragment started, wait for more
return (opcode, payload)
else:
# Control frames
return (opcode, payload)
async def _read_exact(self, n):
"""Read exactly n bytes from internal buffer.
Waits for data via the callback-fed buffer instead of StreamReader.
"""
while len(self._buffer) < n:
if self._eof:
return None
self._data_event.clear()
# Critical: check buffer AGAIN after clearing to avoid race
# condition where data arrives between clear() and wait()
if len(self._buffer) >= n:
break
await self._data_event.wait()
if self._eof and len(self._buffer) < n:
return None
data = bytes(self._buffer[:n])
del self._buffer[:n]
return data
def _unmask(self, payload, masking_key):
"""Unmask WebSocket payload data."""
if not payload:
return payload
# XOR each byte with corresponding mask byte
return bytes(b ^ masking_key[i % 4] for i, b in enumerate(payload))
async def _handle_close(self, payload):
"""Handle incoming close frame."""
if len(payload) >= 2:
self.close_code = struct.unpack("!H", payload[:2])[0]
self.close_reason = payload[2:].decode("utf-8", errors="replace")
else:
self.close_code = CLOSE_NO_STATUS
self.close_reason = ""
self._close_received = True
# Echo close frame back if we haven't already sent one
if not self._close_sent:
await self._send_close(self.close_code, self.close_reason)
self.closed = True
self._close_event.set()
async def _handle_continuation(self, payload): # pylint: disable=unused-argument
"""Handle continuation frame (already processed in _read_frame)."""
# This is called for partial fragments, nothing to do here
async def _send_frame(self, opcode, payload):
"""Send a WebSocket frame.
Server frames are not masked (RFC 6455).
"""
if isinstance(payload, str):
payload = payload.encode("utf-8")
length = len(payload)
frame = bytearray()
# First byte: FIN + opcode
frame.append(0x80 | opcode)
# Second byte: length (no mask bit for server)
if length < 126:
frame.append(length)
elif length < 65536:
frame.append(126)
frame.extend(struct.pack("!H", length))
else:
frame.append(127)
frame.extend(struct.pack("!Q", length))
# Payload
frame.extend(payload)
self.transport.write(bytes(frame))
async def _send_close(self, code, reason=""):
"""Send a close frame."""
if self._close_sent:
return # Already sent
payload = struct.pack("!H", code)
if reason:
payload += reason.encode("utf-8")[:123] # Max 125 bytes total
await self._send_frame(OPCODE_CLOSE, payload)
self._close_sent = True
# If we already received a close, handshake is complete
if self._close_received:
self.closed = True
self._close_event.set()