actualizado 3-sept
This commit is contained in:
@@ -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']
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
+178
@@ -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
@@ -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
|
||||
+2007
File diff suppressed because it is too large
Load Diff
+135
@@ -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
@@ -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()
|
||||
Reference in New Issue
Block a user