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
+8
View File
@@ -0,0 +1,8 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
version_info = (26, 2, 0)
__version__ = ".".join([str(v) for v in version_info])
SERVER = "gunicorn"
SERVER_SOFTWARE = "%s/%s" % (SERVER, __version__)
+10
View File
@@ -0,0 +1,10 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
from gunicorn.app.wsgiapp import run
if __name__ == "__main__":
# see config.py - argparse defaults to basename(argv[0]) == "__main__.py"
# todo: let runpy.run_module take care of argv[0] rewriting
run(prog="gunicorn")
@@ -0,0 +1,3 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
+235
View File
@@ -0,0 +1,235 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import importlib.util
import importlib.machinery
import os
import sys
import traceback
from gunicorn import util
from gunicorn.arbiter import Arbiter
from gunicorn.config import Config, get_default_config_file
from gunicorn import debug
class BaseApplication:
"""
An application interface for configuring and loading
the various necessities for any given web framework.
"""
def __init__(self, usage=None, prog=None):
self.usage = usage
self.cfg = None
self.callable = None
self.prog = prog
self.logger = None
self.do_load_config()
def do_load_config(self):
"""
Loads the configuration
"""
try:
self.load_default_config()
self.load_config()
except Exception as e:
print("\nError: %s" % str(e), file=sys.stderr)
sys.stderr.flush()
sys.exit(1)
def load_default_config(self):
# init configuration
self.cfg = Config(self.usage, prog=self.prog)
def init(self, parser, opts, args):
raise NotImplementedError
def load(self):
raise NotImplementedError
def load_config(self):
"""
This method is used to load the configuration from one or several input(s).
Custom Command line, configuration file.
You have to override this method in your class.
"""
raise NotImplementedError
def reload(self):
self.do_load_config()
if self.cfg.spew:
debug.spew()
def wsgi(self):
if self.callable is None:
self.callable = self.load()
return self.callable
def run(self):
try:
Arbiter(self).run()
except RuntimeError as e:
print("\nError: %s\n" % e, file=sys.stderr)
sys.stderr.flush()
sys.exit(1)
class Application(BaseApplication):
# 'init' and 'load' methods are implemented by WSGIApplication.
# pylint: disable=abstract-method
def chdir(self):
# chdir to the configured path before loading,
# default is the current dir
os.chdir(self.cfg.chdir)
# add the path to sys.path
if self.cfg.chdir not in sys.path:
sys.path.insert(0, self.cfg.chdir)
def get_config_from_filename(self, filename):
if not os.path.exists(filename):
raise RuntimeError("%r doesn't exist" % filename)
ext = os.path.splitext(filename)[1]
try:
module_name = '__config__'
if ext in [".py", ".pyc"]:
spec = importlib.util.spec_from_file_location(module_name, filename)
else:
msg = "configuration file should have a valid Python extension.\n"
util.warn(msg)
loader_ = importlib.machinery.SourceFileLoader(module_name, filename)
spec = importlib.util.spec_from_file_location(module_name, filename, loader=loader_)
mod = importlib.util.module_from_spec(spec)
sys.modules[module_name] = mod
spec.loader.exec_module(mod)
except Exception:
print("Failed to read config file: %s" % filename, file=sys.stderr)
traceback.print_exc()
sys.stderr.flush()
sys.exit(1)
return vars(mod)
def get_config_from_module_name(self, module_name):
return vars(importlib.import_module(module_name))
def load_config_from_module_name_or_filename(self, location):
"""
Loads the configuration file: the file is a python file, otherwise raise an RuntimeError
Exception or stop the process if the configuration file contains a syntax error.
"""
if location.startswith("python:"):
module_name = location[len("python:"):]
cfg = self.get_config_from_module_name(module_name)
else:
if location.startswith("file:"):
filename = location[len("file:"):]
else:
filename = location
cfg = self.get_config_from_filename(filename)
for k, v in cfg.items():
# Ignore unknown names
if k not in self.cfg.settings:
continue
try:
self.cfg.set(k.lower(), v)
except Exception:
print("Invalid value for %s: %s\n" % (k, v), file=sys.stderr)
sys.stderr.flush()
raise
return cfg
def load_config_from_file(self, filename):
return self.load_config_from_module_name_or_filename(location=filename)
def load_config(self):
# parse console args
parser = self.cfg.parser()
args = parser.parse_args()
# optional settings from apps
cfg = self.init(parser, args, args.args)
# set up import paths and follow symlinks
self.chdir()
# Load up the any app specific configuration
if cfg:
for k, v in cfg.items():
self.cfg.set(k.lower(), v)
env_args = parser.parse_args(self.cfg.get_cmd_args_from_env())
if args.config:
self.load_config_from_file(args.config)
elif env_args.config:
self.load_config_from_file(env_args.config)
else:
default_config = get_default_config_file()
if default_config is not None:
self.load_config_from_file(default_config)
# Load up environment configuration
for k, v in vars(env_args).items():
if v is None:
continue
if k == "args":
continue
self.cfg.set(k.lower(), v)
# Lastly, update the configuration with any command line settings.
for k, v in vars(args).items():
if v is None:
continue
if k == "args":
continue
self.cfg.set(k.lower(), v)
# current directory might be changed by the config now
# set up import paths and follow symlinks
self.chdir()
def run(self):
if self.cfg.print_config:
print(self.cfg)
if self.cfg.print_config or self.cfg.check_config:
try:
self.load()
except Exception:
msg = "\nError while loading the application:\n"
print(msg, file=sys.stderr)
traceback.print_exc()
sys.stderr.flush()
sys.exit(1)
sys.exit(0)
if self.cfg.spew:
debug.spew()
if self.cfg.daemon:
if os.environ.get('NOTIFY_SOCKET'):
msg = "Warning: you shouldn't specify `daemon = True`" \
" when launching by systemd with `Type = notify`"
print(msg, file=sys.stderr, flush=True)
util.daemonize(self.cfg.enable_stdio_inheritance)
# set python paths
if self.cfg.pythonpath:
paths = self.cfg.pythonpath.split(",")
for path in paths:
pythonpath = os.path.abspath(path)
if pythonpath not in sys.path:
sys.path.insert(0, pythonpath)
super().run()
@@ -0,0 +1,74 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import configparser
import os
from paste.deploy import loadapp
from gunicorn.app.wsgiapp import WSGIApplication
from gunicorn.config import get_default_config_file
def get_wsgi_app(config_uri, name=None, defaults=None):
if ':' not in config_uri:
config_uri = "config:%s" % config_uri
return loadapp(
config_uri,
name=name,
relative_to=os.getcwd(),
global_conf=defaults,
)
def has_logging_config(config_file):
parser = configparser.ConfigParser()
parser.read([config_file])
return parser.has_section('loggers')
def serve(app, global_conf, **local_conf):
"""\
A Paste Deployment server runner.
Example configuration:
[server:main]
use = egg:gunicorn#main
host = 127.0.0.1
port = 5000
"""
config_file = global_conf['__file__']
gunicorn_config_file = local_conf.pop('config', None)
host = local_conf.pop('host', '')
port = local_conf.pop('port', '')
if host and port:
local_conf['bind'] = '%s:%s' % (host, port)
elif host:
local_conf['bind'] = host.split(',')
class PasterServerApplication(WSGIApplication):
def load_config(self):
self.cfg.set("default_proc_name", config_file)
if has_logging_config(config_file):
self.cfg.set("logconfig", config_file)
if gunicorn_config_file:
self.load_config_from_file(gunicorn_config_file)
else:
default_gunicorn_config_file = get_default_config_file()
if default_gunicorn_config_file is not None:
self.load_config_from_file(default_gunicorn_config_file)
for k, v in local_conf.items():
if v is not None:
self.cfg.set(k.lower(), v)
def load(self):
return app
PasterServerApplication().run()
+70
View File
@@ -0,0 +1,70 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import os
from gunicorn.errors import ConfigError
from gunicorn.app.base import Application
from gunicorn import util
class WSGIApplication(Application):
def init(self, parser, opts, args):
self.app_uri = None
if opts.paste:
from .pasterapp import has_logging_config
config_uri = os.path.abspath(opts.paste)
config_file = config_uri.split('#')[0]
if not os.path.exists(config_file):
raise ConfigError("%r not found" % config_file)
self.cfg.set("default_proc_name", config_file)
self.app_uri = config_uri
if has_logging_config(config_file):
self.cfg.set("logconfig", config_file)
return
if len(args) > 0:
self.cfg.set("default_proc_name", args[0])
self.app_uri = args[0]
def load_config(self):
super().load_config()
if self.app_uri is None:
if self.cfg.wsgi_app is not None:
self.app_uri = self.cfg.wsgi_app
else:
raise ConfigError("No application module specified.")
def load_wsgiapp(self):
return util.import_app(self.app_uri)
def load_pasteapp(self):
from .pasterapp import get_wsgi_app
return get_wsgi_app(self.app_uri, defaults=self.cfg.paste_global_conf)
def load(self):
if self.cfg.paste is not None:
return self.load_pasteapp()
else:
return self.load_wsgiapp()
def run(prog=None):
"""\
The ``gunicorn`` command line runner for launching Gunicorn with
generic WSGI applications.
"""
from gunicorn.app.wsgiapp import WSGIApplication
WSGIApplication("%(prog)s [OPTIONS] [APP_MODULE]", prog=prog).run()
if __name__ == '__main__':
run()
File diff suppressed because it is too large Load Diff
@@ -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()
File diff suppressed because it is too large Load Diff
+16
View File
@@ -0,0 +1,16 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Gunicorn Control Interface
Provides a control socket server for runtime management and
a CLI client (gunicornc) for interacting with running Gunicorn instances.
"""
from gunicorn.ctl.server import ControlSocketServer
from gunicorn.ctl.client import ControlClient
from gunicorn.ctl.protocol import ControlProtocol
__all__ = ['ControlSocketServer', 'ControlClient', 'ControlProtocol']
+450
View File
@@ -0,0 +1,450 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
gunicornc - Gunicorn control interface CLI
Interactive and single-command modes for controlling Gunicorn instances.
"""
import argparse
import json
import os
import sys
from gunicorn.config import _get_default_control_socket
from gunicorn.ctl.client import ControlClient, ControlClientError, parse_command
def format_workers(data: dict) -> str:
"""Format workers output for display."""
workers = data.get("workers", [])
if not workers:
return "No workers running"
lines = []
lines.append(f"{'PID':<10} {'AGE':<6} {'BOOTED':<8} {'LAST_BEAT'}")
lines.append("-" * 40)
for w in workers:
pid = w.get("pid", "?")
age = w.get("age", "?")
booted = "yes" if w.get("booted") else "no"
hb = w.get("last_heartbeat")
hb_str = f"{hb}s ago" if hb is not None else "n/a"
lines.append(f"{pid:<10} {age:<6} {booted:<8} {hb_str}")
lines.append("")
lines.append(f"Total: {data.get('count', len(workers))} workers")
return "\n".join(lines)
def format_dirty(data: dict) -> str:
"""Format dirty workers output for display."""
if not data.get("enabled"):
return "Dirty arbiter not running"
lines = []
lines.append(f"Dirty arbiter PID: {data.get('pid')}")
lines.append("")
workers = data.get("workers", [])
if workers:
lines.append("DIRTY WORKERS:")
lines.append(f"{'PID':<10} {'AGE':<6} {'APPS':<30} {'LAST_BEAT'}")
lines.append("-" * 60)
for w in workers:
pid = w.get("pid", "?")
age = w.get("age", "?")
apps = ", ".join(w.get("apps", []))[:30]
hb = w.get("last_heartbeat")
hb_str = f"{hb}s ago" if hb is not None else "n/a"
lines.append(f"{pid:<10} {age:<6} {apps:<30} {hb_str}")
lines.append("")
apps = data.get("apps", [])
if apps:
lines.append("DIRTY APPS:")
lines.append(f"{'APP':<30} {'WORKERS':<10} {'LIMIT'}")
lines.append("-" * 50)
for app in apps:
path = app.get("import_path", "?")[:30]
current = app.get("current_workers", 0)
limit = app.get("worker_count")
limit_str = str(limit) if limit is not None else "none"
lines.append(f"{path:<30} {current:<10} {limit_str}")
return "\n".join(lines)
def format_stats(data: dict) -> str:
"""Format stats output for display."""
lines = []
uptime = data.get("uptime")
if uptime:
hours = int(uptime // 3600)
minutes = int((uptime % 3600) // 60)
seconds = int(uptime % 60)
if hours:
uptime_str = f"{hours}h {minutes}m {seconds}s"
elif minutes:
uptime_str = f"{minutes}m {seconds}s"
else:
uptime_str = f"{seconds}s"
else:
uptime_str = "unknown"
lines.append(f"Uptime: {uptime_str}")
lines.append(f"PID: {data.get('pid', 'unknown')}")
lines.append(f"Workers current: {data.get('workers_current', 0)}")
lines.append(f"Workers target: {data.get('workers_target', 0)}")
lines.append(f"Workers spawned: {data.get('workers_spawned', 0)}")
lines.append(f"Workers killed: {data.get('workers_killed', 0)}")
lines.append(f"Reloads: {data.get('reloads', 0)}")
dirty_pid = data.get("dirty_arbiter_pid")
if dirty_pid:
lines.append(f"Dirty arbiter: {dirty_pid}")
return "\n".join(lines)
def format_listeners(data: dict) -> str:
"""Format listeners output for display."""
listeners = data.get("listeners", [])
if not listeners:
return "No listeners bound"
lines = []
lines.append(f"{'ADDRESS':<40} {'TYPE':<8} {'FD'}")
lines.append("-" * 55)
for lnr in listeners:
addr = lnr.get("address", "?")
ltype = lnr.get("type", "?")
fd = lnr.get("fd", "?")
lines.append(f"{addr:<40} {ltype:<8} {fd}")
lines.append("")
lines.append(f"Total: {data.get('count', len(listeners))} listeners")
return "\n".join(lines)
def format_config(data: dict) -> str:
"""Format config output for display."""
lines = []
# Sort keys for consistent output
for key in sorted(data.keys()):
value = data[key]
if isinstance(value, list):
value = ", ".join(str(v) for v in value)
lines.append(f"{key}: {value}")
return "\n".join(lines)
def format_help(data: dict) -> str:
"""Format help output for display."""
commands = data.get("commands", {})
lines = []
lines.append("Available commands:")
lines.append("")
# Find max command length for alignment
max_len = max(len(cmd) for cmd in commands.keys()) if commands else 0
for cmd, desc in sorted(commands.items()):
lines.append(f" {cmd:<{max_len + 2}} {desc}")
return "\n".join(lines)
def format_all(data: dict) -> str:
"""Format show all output for display."""
lines = []
# Arbiter
arbiter = data.get("arbiter", {})
lines.append("ARBITER (master)")
lines.append(f" PID: {arbiter.get('pid', '?')}")
lines.append("")
# Web workers
web_workers = data.get("web_workers", [])
lines.append(f"WEB WORKERS ({data.get('web_worker_count', 0)})")
if web_workers:
lines.append(f" {'PID':<10} {'AGE':<6} {'BOOTED':<8} {'LAST_BEAT'}")
lines.append(f" {'-' * 38}")
for w in web_workers:
pid = w.get("pid", "?")
age = w.get("age", "?")
booted = "yes" if w.get("booted") else "no"
hb = w.get("last_heartbeat")
hb_str = f"{hb}s ago" if hb is not None else "n/a"
lines.append(f" {pid:<10} {age:<6} {booted:<8} {hb_str}")
else:
lines.append(" (none)")
lines.append("")
# Dirty arbiter
dirty_arbiter = data.get("dirty_arbiter")
if dirty_arbiter:
lines.append("DIRTY ARBITER")
lines.append(f" PID: {dirty_arbiter.get('pid', '?')}")
lines.append("")
# Dirty workers
dirty_workers = data.get("dirty_workers", [])
lines.append(f"DIRTY WORKERS ({data.get('dirty_worker_count', 0)})")
if dirty_workers:
lines.append(f" {'PID':<10} {'AGE':<6} {'APPS'}")
lines.append(f" {'-' * 50}")
for w in dirty_workers:
pid = w.get("pid", "?")
age = w.get("age", "?")
apps = w.get("apps", [])
# Show each app on its own line if multiple
if apps:
first_app = apps[0].split(":")[-1] # Just the class name
lines.append(f" {pid:<10} {age:<6} {first_app}")
for app in apps[1:]:
app_name = app.split(":")[-1]
lines.append(f" {'':<10} {'':<6} {app_name}")
else:
lines.append(f" {pid:<10} {age:<6} (no apps)")
else:
lines.append(" (none)")
else:
lines.append("DIRTY ARBITER")
lines.append(" (not running)")
return "\n".join(lines)
def format_response(command: str, data: dict) -> str: # pylint: disable=too-many-return-statements
"""
Format response data based on command.
Args:
command: Original command string
data: Response data dictionary
Returns:
Formatted string for display
"""
cmd_lower = command.lower().strip()
# Route to specific formatters
if cmd_lower == "show all":
return format_all(data)
elif cmd_lower == "show workers":
return format_workers(data)
elif cmd_lower == "show dirty":
return format_dirty(data)
elif cmd_lower == "show stats":
return format_stats(data)
elif cmd_lower == "show listeners":
return format_listeners(data)
elif cmd_lower == "show config":
return format_config(data)
elif cmd_lower == "help":
return format_help(data)
else:
# Generic JSON output for other commands
if data:
return json.dumps(data, indent=2)
return "OK"
def run_command(socket_path: str, command: str, json_output: bool = False) -> int:
"""
Execute single command and exit.
Args:
socket_path: Path to control socket
command: Command to execute
json_output: If True, output raw JSON
Returns:
Exit code (0 for success, 1 for error)
"""
try:
with ControlClient(socket_path) as client:
cmd, args = parse_command(command)
full_command = f"{cmd} {' '.join(args)}".strip() if args else cmd
result = client.send_command(full_command)
if json_output:
print(json.dumps(result, indent=2))
else:
output = format_response(cmd, result)
print(output)
return 0
except ControlClientError as e:
print(f"Error: {e}", file=sys.stderr)
return 1
except KeyboardInterrupt:
return 130
def run_interactive(socket_path: str, json_output: bool = False) -> int:
"""
Run interactive CLI with readline support.
Args:
socket_path: Path to control socket
json_output: If True, output raw JSON
Returns:
Exit code
"""
try:
import readline # noqa: F401 - imported for side effects
has_readline = True
except ImportError:
has_readline = False
try:
client = ControlClient(socket_path)
client.connect()
except ControlClientError as e:
print(f"Error: {e}", file=sys.stderr)
return 1
print(f"Connected to {socket_path}")
print("Type 'help' for available commands, 'quit' to exit.")
print()
# Set up readline history
history_file = os.path.expanduser("~/.gunicornc_history")
if has_readline:
try:
readline.read_history_file(history_file)
except FileNotFoundError:
pass
exit_code = 0
try:
while True:
try:
line = input("gunicorn> ").strip()
except EOFError:
print()
break
if not line:
continue
if line.lower() in ('quit', 'exit', 'q'):
break
try:
cmd, args = parse_command(line)
full_command = f"{cmd} {' '.join(args)}".strip() if args else cmd
result = client.send_command(full_command)
if json_output:
print(json.dumps(result, indent=2))
else:
output = format_response(cmd, result)
print(output)
except ControlClientError as e:
print(f"Error: {e}")
# Try to reconnect
try:
client.close()
client.connect()
except ControlClientError:
print("Connection lost. Exiting.")
exit_code = 1
break
print()
except KeyboardInterrupt:
print()
exit_code = 130
finally:
client.close()
if has_readline:
try:
readline.write_history_file(history_file)
except Exception:
pass
return exit_code
def main():
"""Main entry point for gunicornc CLI."""
parser = argparse.ArgumentParser(
description='Gunicorn control interface',
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
gunicornc # Interactive mode (default socket)
gunicornc -s /tmp/myapp.ctl # Interactive mode with custom socket
gunicornc -c "show workers" # Single command mode
gunicornc -c "worker add 2" # Add 2 workers
gunicornc -c "show stats" -j # Output stats as JSON
"""
)
parser.add_argument(
'-s', '--socket',
default=_get_default_control_socket(),
help='Control socket path (default: auto-detected based on XDG_RUNTIME_DIR or ~/.gunicorn/)'
)
parser.add_argument(
'-c', '--command',
help='Execute single command and exit'
)
parser.add_argument(
'-j', '--json',
action='store_true',
help='Output raw JSON (for scripting)'
)
parser.add_argument(
'-v', '--version',
action='store_true',
help='Show version and exit'
)
args = parser.parse_args()
if args.version:
from gunicorn import __version__
print(f"gunicornc (gunicorn {__version__})")
return 0
socket_path = args.socket
# Make relative paths absolute from cwd
if not os.path.isabs(socket_path):
socket_path = os.path.join(os.getcwd(), socket_path)
if args.command:
return run_command(socket_path, args.command, args.json)
else:
return run_interactive(socket_path, args.json)
if __name__ == '__main__':
sys.exit(main())
+139
View File
@@ -0,0 +1,139 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Control Socket Client
Client library for connecting to gunicorn control socket.
"""
import shlex
import socket
from gunicorn.ctl.protocol import (
ControlProtocol,
make_request,
)
class ControlClientError(Exception):
"""Control client error."""
class ControlClient:
"""
Client for connecting to gunicorn control socket.
Can be used as a context manager:
with ControlClient('/path/to/gunicorn.ctl') as client:
result = client.send_command('show workers')
"""
def __init__(self, socket_path: str, timeout: float = 30.0):
"""
Initialize control client.
Args:
socket_path: Path to the Unix socket
timeout: Socket timeout in seconds (default 30)
"""
self.socket_path = socket_path
self.timeout = timeout
self._sock = None
self._request_id = 0
def connect(self):
"""
Connect to control socket.
Raises:
ControlClientError: If connection fails
"""
if self._sock:
return
try:
self._sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
self._sock.settimeout(self.timeout)
self._sock.connect(self.socket_path)
except socket.error as e:
self._sock = None
raise ControlClientError(f"Failed to connect to {self.socket_path}: {e}")
def close(self):
"""Close connection."""
if self._sock:
try:
self._sock.close()
except Exception:
pass
self._sock = None
def send_command(self, command: str, args: list = None) -> dict:
"""
Send command and wait for response.
Args:
command: Command string (e.g., "show workers")
args: Optional additional arguments
Returns:
Response data dictionary
Raises:
ControlClientError: If communication fails
"""
if not self._sock:
self.connect()
self._request_id += 1
request = make_request(self._request_id, command, args)
try:
ControlProtocol.write_message(self._sock, request)
response = ControlProtocol.read_message(self._sock)
except Exception as e:
self.close()
raise ControlClientError(f"Communication error: {e}")
if response.get("status") == "error":
raise ControlClientError(response.get("error", "Unknown error"))
return response.get("data", {})
def __enter__(self):
self.connect()
return self
def __exit__(self, *args):
self.close()
def parse_command(line: str) -> tuple:
"""
Parse a command line into command and args.
Args:
line: Command line string
Returns:
Tuple of (command_string, args_list)
"""
parts = shlex.split(line)
if not parts:
return "", []
# Find where numeric/value args start
command_parts = []
args = []
for part in parts:
# If we haven't hit args yet and this looks like a command word
if not args and not part.isdigit() and not part.startswith('-'):
command_parts.append(part)
else:
args.append(part)
return " ".join(command_parts), args
+585
View File
@@ -0,0 +1,585 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Control Interface Command Handlers
Provides handlers for all control commands with access to arbiter state.
"""
import os
import signal
import socket
import time
class CommandHandlers:
"""
Command handlers with access to arbiter state.
All handler methods return dictionaries that will be sent
as the response data.
"""
def __init__(self, arbiter):
"""
Initialize handlers with arbiter reference.
Args:
arbiter: The Gunicorn arbiter instance
"""
self.arbiter = arbiter
def show_workers(self) -> dict:
"""
Return list of HTTP workers.
Returns:
Dictionary with workers list containing:
- pid: Worker process ID
- age: Worker age (spawn order)
- requests: Number of requests handled (if available)
- booted: Whether worker has finished booting
- last_heartbeat: Seconds since last heartbeat
"""
workers = []
now = time.monotonic()
for pid, worker in self.arbiter.WORKERS.items():
try:
last_update = worker.tmp.last_update()
last_heartbeat = round(now - last_update, 2)
except (OSError, ValueError):
last_heartbeat = None
workers.append({
"pid": pid,
"age": worker.age,
"booted": worker.booted,
"aborted": worker.aborted,
"last_heartbeat": last_heartbeat,
})
# Sort by age (oldest first)
workers.sort(key=lambda w: w["age"])
return {"workers": workers, "count": len(workers)}
def show_dirty(self) -> dict:
"""
Return dirty workers and apps information.
Returns:
Dictionary with:
- enabled: Whether dirty arbiter is running
- pid: Dirty arbiter PID
- workers: List of dirty worker info
- apps: List of dirty app specs
"""
if not self.arbiter.dirty_arbiter_pid:
return {
"enabled": False,
"pid": None,
"workers": [],
"apps": [],
}
# Get dirty arbiter reference if available
dirty_arbiter = getattr(self.arbiter, 'dirty_arbiter', None)
workers = []
apps = []
if dirty_arbiter and hasattr(dirty_arbiter, 'workers'):
now = time.monotonic()
for pid, worker in dirty_arbiter.workers.items():
try:
last_update = worker.tmp.last_update()
last_heartbeat = round(now - last_update, 2)
except (OSError, ValueError, AttributeError):
last_heartbeat = None
workers.append({
"pid": pid,
"age": worker.age,
"apps": getattr(worker, 'app_paths', []),
"booted": getattr(worker, 'booted', False),
"last_heartbeat": last_heartbeat,
})
# Get app specs
if hasattr(dirty_arbiter, 'app_specs'):
for path, spec in dirty_arbiter.app_specs.items():
worker_pids = list(dirty_arbiter.app_worker_map.get(path, []))
apps.append({
"import_path": path,
"worker_count": spec.get('worker_count'),
"current_workers": len(worker_pids),
"worker_pids": worker_pids,
})
return {
"enabled": True,
"pid": self.arbiter.dirty_arbiter_pid,
"workers": workers,
"apps": apps,
}
def show_config(self) -> dict:
"""
Return current effective configuration.
Returns:
Dictionary of configuration values
"""
cfg = self.arbiter.cfg
config = {}
# Get commonly needed config values
config_keys = [
'bind', 'workers', 'worker_class', 'threads', 'timeout',
'graceful_timeout', 'keepalive', 'max_requests',
'max_requests_jitter', 'worker_connections', 'preload_app',
'daemon', 'pidfile', 'proc_name', 'reload',
'dirty_workers', 'dirty_apps', 'dirty_timeout',
'control_socket', 'control_socket_disable',
]
for key in config_keys:
try:
value = getattr(cfg, key)
# Convert non-serializable types
if callable(value):
value = str(value)
elif hasattr(value, '__class__') and not isinstance(
value, (str, int, float, bool, list, dict, type(None))):
value = str(value)
config[key] = value
except AttributeError:
pass
return config
def show_stats(self) -> dict:
"""
Return server statistics.
Returns:
Dictionary with:
- uptime: Seconds since arbiter started
- pid: Arbiter PID
- workers_current: Current number of workers
- workers_spawned: Total workers spawned
- workers_killed: Total workers killed (if tracked)
- reloads: Number of reloads (if tracked)
"""
stats = getattr(self.arbiter, '_stats', {})
start_time = stats.get('start_time')
uptime = None
if start_time:
uptime = round(time.time() - start_time, 2)
return {
"uptime": uptime,
"pid": self.arbiter.pid,
"workers_current": len(self.arbiter.WORKERS),
"workers_target": self.arbiter.num_workers,
"workers_spawned": stats.get('workers_spawned', 0),
"workers_killed": stats.get('workers_killed', 0),
"reloads": stats.get('reloads', 0),
"dirty_arbiter_pid": self.arbiter.dirty_arbiter_pid or None,
}
def show_listeners(self) -> dict:
"""
Return bound socket information.
Returns:
Dictionary with listeners list
"""
listeners = []
for lnr in self.arbiter.LISTENERS:
addr = str(lnr)
listener_info = {
"address": addr,
"fd": lnr.fileno(),
}
# Try to get socket family
try:
sock = lnr.sock
if sock.family == socket.AF_UNIX:
listener_info["type"] = "unix"
elif sock.family == socket.AF_INET:
listener_info["type"] = "tcp"
elif sock.family == socket.AF_INET6:
listener_info["type"] = "tcp6"
except Exception:
listener_info["type"] = "unknown"
listeners.append(listener_info)
return {"listeners": listeners, "count": len(listeners)}
def worker_add(self, count: int = 1) -> dict:
"""
Increase worker count.
Args:
count: Number of workers to add (default 1)
Returns:
Dictionary with added count and new total
"""
count = max(1, int(count))
old_count = self.arbiter.num_workers
self.arbiter.num_workers += count
# Wake up the arbiter to spawn workers
self.arbiter.wakeup()
return {
"added": count,
"previous": old_count,
"total": self.arbiter.num_workers,
}
def worker_remove(self, count: int = 1) -> dict:
"""
Decrease worker count.
Args:
count: Number of workers to remove (default 1)
Returns:
Dictionary with removed count and new total
"""
count = max(1, int(count))
old_count = self.arbiter.num_workers
# Don't go below 1 worker
new_count = max(1, old_count - count)
actual_removed = old_count - new_count
self.arbiter.num_workers = new_count
# Wake up the arbiter to kill excess workers
self.arbiter.wakeup()
return {
"removed": actual_removed,
"previous": old_count,
"total": new_count,
}
def worker_kill(self, pid: int) -> dict:
"""
Gracefully terminate a specific worker.
Args:
pid: Worker process ID
Returns:
Dictionary with killed PID or error
"""
pid = int(pid)
if pid not in self.arbiter.WORKERS:
return {
"success": False,
"error": f"Worker {pid} not found",
}
try:
os.kill(pid, signal.SIGTERM)
return {
"success": True,
"killed": pid,
}
except OSError as e:
return {
"success": False,
"error": str(e),
}
def dirty_add(self, count: int = 1) -> dict:
"""
Spawn additional dirty workers.
Sends a MANAGE message to the dirty arbiter to spawn workers.
Args:
count: Number of dirty workers to add (default 1)
Returns:
Dictionary with added count or error
"""
if not self.arbiter.dirty_arbiter_pid:
return {
"success": False,
"error": "Dirty arbiter not running",
}
count = max(1, int(count))
return self._send_manage_message("add", count)
def dirty_remove(self, count: int = 1) -> dict:
"""
Remove dirty workers.
Sends a MANAGE message to the dirty arbiter to remove workers.
Args:
count: Number of dirty workers to remove (default 1)
Returns:
Dictionary with removed count or error
"""
if not self.arbiter.dirty_arbiter_pid:
return {
"success": False,
"error": "Dirty arbiter not running",
}
count = max(1, int(count))
return self._send_manage_message("remove", count)
def _send_manage_message(self, operation: str, count: int) -> dict:
"""
Send a worker management message to the dirty arbiter.
Args:
operation: "add" or "remove"
count: Number of workers to add/remove
Returns:
Dictionary with result or error
"""
# Get socket path from arbiter object or environment
dirty_socket_path = None
if hasattr(self.arbiter, 'dirty_arbiter') and self.arbiter.dirty_arbiter:
dirty_socket_path = getattr(
self.arbiter.dirty_arbiter, 'socket_path', None
)
if not dirty_socket_path:
dirty_socket_path = os.environ.get('GUNICORN_DIRTY_SOCKET')
if not dirty_socket_path:
return {
"success": False,
"error": "Cannot find dirty arbiter socket path",
}
try:
from gunicorn.dirty.protocol import (
DirtyProtocol, MANAGE_OP_ADD, MANAGE_OP_REMOVE
)
op = MANAGE_OP_ADD if operation == "add" else MANAGE_OP_REMOVE
sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
sock.settimeout(10.0)
sock.connect(dirty_socket_path)
# Send manage request
request = {
"type": DirtyProtocol.MSG_TYPE_MANAGE,
"id": 1,
"op": op,
"count": count,
}
DirtyProtocol.write_message(sock, request)
# Read response
response = DirtyProtocol.read_message(sock)
sock.close()
if response.get("type") == DirtyProtocol.MSG_TYPE_RESPONSE:
return response.get("result", {"success": True})
elif response.get("type") == DirtyProtocol.MSG_TYPE_ERROR:
error = response.get("error", {})
return {
"success": False,
"error": error.get("message", str(error)),
}
else:
return {
"success": False,
"error": f"Unexpected response type: {response.get('type')}",
}
except Exception as e:
return {
"success": False,
"error": str(e),
}
def reload(self) -> dict:
"""
Trigger graceful reload (equivalent to SIGHUP).
Returns:
Dictionary with status
"""
# Send HUP to self to trigger reload
os.kill(self.arbiter.pid, signal.SIGHUP)
return {"status": "reloading"}
def reopen(self) -> dict:
"""
Reopen log files (equivalent to SIGUSR1).
Returns:
Dictionary with status
"""
os.kill(self.arbiter.pid, signal.SIGUSR1)
return {"status": "reopening"}
def shutdown(self, mode: str = "graceful") -> dict:
"""
Initiate shutdown.
Args:
mode: "graceful" (SIGTERM) or "quick" (SIGINT)
Returns:
Dictionary with status
"""
if mode == "quick":
os.kill(self.arbiter.pid, signal.SIGINT)
else:
os.kill(self.arbiter.pid, signal.SIGTERM)
return {"status": "shutting_down", "mode": mode}
def show_all(self) -> dict:
"""
Return overview of all processes (arbiter, web workers, dirty arbiter, dirty workers).
Returns:
Dictionary with complete process hierarchy
"""
now = time.monotonic()
# Arbiter info
arbiter_info = {
"pid": self.arbiter.pid,
"type": "arbiter",
"role": "master",
}
# Web workers (HTTP workers)
web_workers = []
for pid, worker in self.arbiter.WORKERS.items():
try:
last_update = worker.tmp.last_update()
last_heartbeat = round(now - last_update, 2)
except (OSError, ValueError):
last_heartbeat = None
web_workers.append({
"pid": pid,
"type": "web",
"age": worker.age,
"booted": worker.booted,
"last_heartbeat": last_heartbeat,
})
# Sort by age
web_workers.sort(key=lambda w: w["age"])
# Dirty arbiter info (runs in separate process)
dirty_arbiter_info = None
dirty_workers = []
if self.arbiter.dirty_arbiter_pid:
dirty_arbiter_info = {
"pid": self.arbiter.dirty_arbiter_pid,
"type": "dirty_arbiter",
"role": "dirty master",
}
# Query dirty arbiter for worker info via its socket
dirty_workers = self._query_dirty_workers()
return {
"arbiter": arbiter_info,
"web_workers": web_workers,
"web_worker_count": len(web_workers),
"dirty_arbiter": dirty_arbiter_info,
"dirty_workers": dirty_workers,
"dirty_worker_count": len(dirty_workers),
}
def _query_dirty_workers(self) -> list:
"""
Query the dirty arbiter for worker information.
Connects to the dirty arbiter socket and sends a status request.
Returns:
List of dirty worker info dicts, or empty list on error
"""
# Get socket path from arbiter object or environment
dirty_socket_path = None
if hasattr(self.arbiter, 'dirty_arbiter') and self.arbiter.dirty_arbiter:
dirty_socket_path = getattr(self.arbiter.dirty_arbiter, 'socket_path', None)
if not dirty_socket_path:
dirty_socket_path = os.environ.get('GUNICORN_DIRTY_SOCKET')
if not dirty_socket_path:
return []
try:
from gunicorn.dirty.protocol import DirtyProtocol
sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
sock.settimeout(2.0)
sock.connect(dirty_socket_path)
# Send status request
request = {
"type": DirtyProtocol.MSG_TYPE_STATUS,
"id": "ctl-status-1",
}
DirtyProtocol.write_message(sock, request)
# Read response
response = DirtyProtocol.read_message(sock)
sock.close()
if response.get("type") == DirtyProtocol.MSG_TYPE_RESPONSE:
result = response.get("result", {})
return result.get("workers", [])
except Exception:
pass
return []
def help(self) -> dict:
"""
Return list of available commands.
Returns:
Dictionary with commands and descriptions
"""
commands = {
"show all": "Show all processes (arbiter, web workers, dirty workers)",
"show workers": "List HTTP workers with their status",
"show dirty": "List dirty workers and apps",
"show config": "Show current effective configuration",
"show stats": "Show server statistics",
"show listeners": "Show bound sockets",
"worker add [N]": "Spawn N workers (default 1)",
"worker remove [N]": "Remove N workers (default 1)",
"worker kill <PID>": "Gracefully terminate specific worker",
"dirty add [N]": "Spawn N dirty workers (default 1)",
"dirty remove [N]": "Remove N dirty workers (default 1)",
"reload": "Graceful reload (HUP)",
"reopen": "Reopen log files (USR1)",
"shutdown [graceful|quick]": "Shutdown server (TERM/INT)",
"help": "Show this help message",
}
return {"commands": commands}
+224
View File
@@ -0,0 +1,224 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Control Socket Protocol
JSON-based protocol with length-prefixed framing for the control interface.
Message Format:
+----------------+------------------+
| Length (4B BE) | JSON Payload |
+----------------+------------------+
Request Format:
{"id": 1, "command": "show", "args": ["workers"]}
Response Format:
{"id": 1, "status": "ok", "data": {...}}
{"id": 1, "status": "error", "error": "message"}
"""
import json
import struct
class ProtocolError(Exception):
"""Protocol-level error."""
class ControlProtocol:
"""
Protocol implementation for control socket communication.
Uses 4-byte big-endian length prefix followed by JSON payload.
"""
# Maximum message size (16 MB)
MAX_MESSAGE_SIZE = 16 * 1024 * 1024
@staticmethod
def encode_message(data: dict) -> bytes:
"""
Encode a message for transmission.
Args:
data: Dictionary to encode
Returns:
Length-prefixed JSON bytes
"""
payload = json.dumps(data).encode('utf-8')
length = struct.pack('>I', len(payload))
return length + payload
@staticmethod
def decode_message(data: bytes) -> dict:
"""
Decode a message from bytes.
Args:
data: Raw bytes (length prefix + JSON payload)
Returns:
Decoded dictionary
"""
if len(data) < 4:
raise ProtocolError("Message too short")
length = struct.unpack('>I', data[:4])[0]
if len(data) < 4 + length:
raise ProtocolError("Incomplete message")
payload = data[4:4 + length]
return json.loads(payload.decode('utf-8'))
@staticmethod
def read_message(sock) -> dict:
"""
Read one message from a socket.
Args:
sock: Socket to read from
Returns:
Decoded message dictionary
Raises:
ProtocolError: If message is malformed
ConnectionError: If connection is closed
"""
# Read length prefix
length_data = b''
while len(length_data) < 4:
chunk = sock.recv(4 - len(length_data))
if not chunk:
if not length_data:
raise ConnectionError("Connection closed")
raise ProtocolError("Incomplete length prefix")
length_data += chunk
length = struct.unpack('>I', length_data)[0]
if length > ControlProtocol.MAX_MESSAGE_SIZE:
raise ProtocolError(f"Message too large: {length}")
# Read payload
payload_data = b''
while len(payload_data) < length:
chunk = sock.recv(min(length - len(payload_data), 65536))
if not chunk:
raise ProtocolError("Incomplete payload")
payload_data += chunk
try:
return json.loads(payload_data.decode('utf-8'))
except json.JSONDecodeError as e:
raise ProtocolError(f"Invalid JSON: {e}")
@staticmethod
def write_message(sock, data: dict):
"""
Write one message to a socket.
Args:
sock: Socket to write to
data: Message dictionary to send
"""
message = ControlProtocol.encode_message(data)
sock.sendall(message)
@staticmethod
async def read_message_async(reader) -> dict:
"""
Read one message from an async reader.
Args:
reader: asyncio StreamReader
Returns:
Decoded message dictionary
"""
# Read length prefix
length_data = await reader.readexactly(4)
length = struct.unpack('>I', length_data)[0]
if length > ControlProtocol.MAX_MESSAGE_SIZE:
raise ProtocolError(f"Message too large: {length}")
# Read payload
payload_data = await reader.readexactly(length)
try:
return json.loads(payload_data.decode('utf-8'))
except json.JSONDecodeError as e:
raise ProtocolError(f"Invalid JSON: {e}")
@staticmethod
async def write_message_async(writer, data: dict):
"""
Write one message to an async writer.
Args:
writer: asyncio StreamWriter
data: Message dictionary to send
"""
message = ControlProtocol.encode_message(data)
writer.write(message)
await writer.drain()
def make_request(request_id: int, command: str, args: list = None) -> dict:
"""
Create a request message.
Args:
request_id: Unique request identifier
command: Command name (e.g., "show workers")
args: Optional list of arguments
Returns:
Request dictionary
"""
return {
"id": request_id,
"command": command,
"args": args or [],
}
def make_response(request_id: int, data: dict = None) -> dict:
"""
Create a success response message.
Args:
request_id: Request identifier being responded to
data: Response data
Returns:
Response dictionary
"""
return {
"id": request_id,
"status": "ok",
"data": data or {},
}
def make_error_response(request_id: int, error: str) -> dict:
"""
Create an error response message.
Args:
request_id: Request identifier being responded to
error: Error message
Returns:
Error response dictionary
"""
return {
"id": request_id,
"status": "error",
"error": error,
}
+419
View File
@@ -0,0 +1,419 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Control Socket Server
Runs in the arbiter process and accepts commands via Unix socket.
Uses asyncio in a background thread to handle client connections.
Fork Safety:
This server uses os.register_at_fork() to properly handle fork() calls.
Before fork: the asyncio thread is stopped to prevent lock issues.
After fork in parent: the server is restarted.
After fork in child: references are cleared (workers don't need the control server).
"""
import asyncio
import os
import shlex
import threading
from gunicorn.ctl.handlers import CommandHandlers
from gunicorn.ctl.protocol import (
ControlProtocol,
make_response,
make_error_response,
)
# Module-level tracking of active control server instances for fork handling.
# This is necessary because os.register_at_fork() callbacks are process-level.
_active_servers = set()
_module_state = {"fork_handlers_registered": False}
def _register_fork_handlers():
"""Register fork handlers once at module level."""
if _module_state["fork_handlers_registered"]:
return
_module_state["fork_handlers_registered"] = True
os.register_at_fork(
before=_before_fork,
after_in_parent=_after_fork_parent,
after_in_child=_after_fork_child,
)
def _before_fork():
"""Called before fork() - stop all active control servers."""
for server in list(_active_servers):
server._stop_for_fork()
def _after_fork_parent():
"""Called in parent after fork() - restart all control servers."""
for server in list(_active_servers):
server._restart_after_fork()
def _after_fork_child():
"""Called in child after fork() - cleanup references."""
# In the child process (worker), we don't need the control server.
# Just clear the references without trying to stop anything.
_active_servers.clear()
class ControlSocketServer:
"""
Control socket server running in arbiter process.
The server runs an asyncio event loop in a background thread,
accepting connections and dispatching commands to handlers.
Fork safety is handled via os.register_at_fork() - the server
automatically stops before fork and restarts after in the parent.
"""
def __init__(self, arbiter, socket_path, socket_mode=0o600):
"""
Initialize control socket server.
Args:
arbiter: The Gunicorn arbiter instance
socket_path: Path for the Unix socket
socket_mode: Permission mode for socket (default 0o600)
"""
self.arbiter = arbiter
self.socket_path = socket_path
self.socket_mode = socket_mode
self.handlers = CommandHandlers(arbiter)
self._server = None
self._loop = None
self._thread = None
self._running = False
self._was_running_before_fork = False
# Set by the loop thread once self._loop and self._server are live.
# The stop paths wait on it so a shutdown is never scheduled against a
# not-yet-initialized loop (which would leak the thread + its fds).
self._ready = threading.Event()
# Ensure fork handlers are registered
_register_fork_handlers()
def start(self):
"""Start server in background thread with asyncio event loop."""
if self._running:
return
self._running = True
self._ready = threading.Event()
self._thread = threading.Thread(target=self._run_loop, daemon=True)
self._thread.start()
# Track this server for fork handling
_active_servers.add(self)
def stop(self):
"""Stop server and cleanup socket."""
# Remove from active servers tracking
_active_servers.discard(self)
if not self._running:
return
self._running = False
# Wait until the loop is live so the shutdown is actually delivered.
self._ready.wait(timeout=2.0)
if self._loop is not None:
# Schedule server close in the loop
self._loop.call_soon_threadsafe(self._shutdown)
if self._thread:
self._thread.join(timeout=2.0)
self._thread = None
# Clean up socket file
if os.path.exists(self.socket_path):
try:
os.unlink(self.socket_path)
except OSError:
pass
def _stop_for_fork(self):
"""Stop server before fork (called by fork handler)."""
if not self._running:
self._was_running_before_fork = False
return
self._was_running_before_fork = True
self._running = False
# Wait until the loop is live before scheduling shutdown. Without this,
# a fork that lands while the thread is still starting up skips the
# shutdown, join() times out, and the thread is dropped while still
# holding its selector fd and the unix socket (the #3648 leak).
self._ready.wait(timeout=2.0)
if self._loop is not None:
try:
self._loop.call_soon_threadsafe(self._shutdown)
except RuntimeError:
# Loop may already be closed
pass
if self._thread:
self._thread.join(timeout=2.0)
if self._thread.is_alive() and self.arbiter.log:
self.arbiter.log.warning(
"control socket thread did not stop before fork")
self._thread = None
self._loop = None
self._server = None
def _restart_after_fork(self):
"""Restart server in parent after fork (called by fork handler)."""
if not self._was_running_before_fork:
return
self._was_running_before_fork = False
self._running = True
self._ready = threading.Event()
self._thread = threading.Thread(target=self._run_loop, daemon=True)
self._thread.start()
def _shutdown(self):
"""Shutdown server (called from event loop thread)."""
if self._server:
self._server.close()
def _run_loop(self):
"""Run the asyncio event loop in background thread."""
try:
asyncio.run(self._serve())
except Exception as e:
if self._running and self.arbiter.log:
self.arbiter.log.error("Control server error: %s", e)
finally:
# Release any stop path waiting on readiness, even if the thread
# never became ready (eg the server failed to bind).
self._ready.set()
async def _serve(self):
"""Main async server loop."""
self._loop = asyncio.get_running_loop()
# Create parent directory if needed (for ~/.gunicorn/)
socket_dir = os.path.dirname(self.socket_path)
if socket_dir and not os.path.exists(socket_dir):
os.makedirs(socket_dir, mode=0o700)
# Remove socket if it exists
if os.path.exists(self.socket_path):
os.unlink(self.socket_path)
# Create Unix socket server
self._server = await asyncio.start_unix_server(
self._handle_client,
path=self.socket_path
)
# Set socket permissions
os.chmod(self.socket_path, self.socket_mode)
if self.arbiter.log:
self.arbiter.log.info("Control socket listening at %s",
self.socket_path)
# Loop and server are live; let the stop paths schedule shutdown.
self._ready.set()
try:
async with self._server:
await self._server.serve_forever()
except asyncio.CancelledError:
pass
finally:
if os.path.exists(self.socket_path):
try:
os.unlink(self.socket_path)
except OSError:
pass
async def _handle_client(self, reader, writer):
"""
Handle client connection.
Args:
reader: asyncio StreamReader
writer: asyncio StreamWriter
"""
try:
while self._running:
try:
message = await asyncio.wait_for(
ControlProtocol.read_message_async(reader),
timeout=300.0 # 5 minute idle timeout
)
except asyncio.TimeoutError:
# Client idle too long, close connection
break
except asyncio.IncompleteReadError:
# Client disconnected
break
except Exception:
# Protocol error
break
# Process command
response = await self._dispatch(message)
# Send response
await ControlProtocol.write_message_async(writer, response)
except Exception as e:
if self.arbiter.log:
self.arbiter.log.debug("Control client error: %s", e)
finally:
writer.close()
try:
await writer.wait_closed()
except Exception:
pass
async def _dispatch(self, message: dict) -> dict:
"""
Dispatch command to appropriate handler.
Args:
message: Request message dict
Returns:
Response dictionary
"""
request_id = message.get("id", 0)
command = message.get("command", "").strip()
args = message.get("args", [])
if not command:
return make_error_response(request_id, "Empty command")
try:
# Parse command (e.g., "show workers" or "worker add 2")
parts = shlex.split(command)
if args:
parts.extend(str(a) for a in args)
if not parts:
return make_error_response(request_id, "Empty command")
# Route to handler
result = self._execute_command(parts)
return make_response(request_id, result)
except ValueError as e:
return make_error_response(request_id, f"Invalid argument: {e}")
except Exception as e:
if self.arbiter.log:
self.arbiter.log.exception("Command error")
return make_error_response(request_id, f"Command failed: {e}")
def _execute_command(self, parts: list) -> dict: # pylint: disable=too-many-return-statements
"""
Execute a parsed command.
Args:
parts: Command parts (e.g., ["show", "workers"])
Returns:
Handler result dictionary
"""
if not parts:
raise ValueError("Empty command")
cmd = parts[0].lower()
rest = parts[1:]
# Map commands to handlers
if cmd == "show":
return self._handle_show(rest)
elif cmd == "worker":
return self._handle_worker(rest)
elif cmd == "dirty":
return self._handle_dirty(rest)
elif cmd == "reload":
return self.handlers.reload()
elif cmd == "reopen":
return self.handlers.reopen()
elif cmd == "shutdown":
mode = rest[0] if rest else "graceful"
return self.handlers.shutdown(mode)
elif cmd == "help":
return self.handlers.help()
else:
raise ValueError(f"Unknown command: {cmd}")
def _handle_show(self, args: list) -> dict:
"""Handle 'show' commands."""
if not args:
raise ValueError("Missing show target (all|workers|dirty|config|stats|listeners)")
target = args[0].lower()
if target == "all":
return self.handlers.show_all()
elif target == "workers":
return self.handlers.show_workers()
elif target == "dirty":
return self.handlers.show_dirty()
elif target == "config":
return self.handlers.show_config()
elif target == "stats":
return self.handlers.show_stats()
elif target == "listeners":
return self.handlers.show_listeners()
else:
raise ValueError(f"Unknown show target: {target}")
def _handle_worker(self, args: list) -> dict:
"""Handle 'worker' commands."""
if not args:
raise ValueError("Missing worker action (add|remove|kill)")
action = args[0].lower()
action_args = args[1:]
if action == "add":
count = int(action_args[0]) if action_args else 1
return self.handlers.worker_add(count)
elif action == "remove":
count = int(action_args[0]) if action_args else 1
return self.handlers.worker_remove(count)
elif action == "kill":
if not action_args:
raise ValueError("Missing PID for worker kill")
pid = int(action_args[0])
return self.handlers.worker_kill(pid)
else:
raise ValueError(f"Unknown worker action: {action}")
def _handle_dirty(self, args: list) -> dict:
"""Handle 'dirty' commands."""
if not args:
raise ValueError("Missing dirty action (add|remove)")
action = args[0].lower()
action_args = args[1:]
if action == "add":
count = int(action_args[0]) if action_args else 1
return self.handlers.dirty_add(count)
elif action == "remove":
count = int(action_args[0]) if action_args else 1
return self.handlers.dirty_remove(count)
else:
raise ValueError(f"Unknown dirty action: {action}")
+68
View File
@@ -0,0 +1,68 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""The debug module contains utilities and functions for better
debugging Gunicorn."""
import sys
import linecache
import re
import inspect
__all__ = ['spew', 'unspew']
_token_spliter = re.compile(r'\W+')
class Spew:
def __init__(self, trace_names=None, show_values=True):
self.trace_names = trace_names
self.show_values = show_values
def __call__(self, frame, event, arg):
if event == 'line':
lineno = frame.f_lineno
if '__file__' in frame.f_globals:
filename = frame.f_globals['__file__']
if (filename.endswith('.pyc') or
filename.endswith('.pyo')):
filename = filename[:-1]
name = frame.f_globals['__name__']
line = linecache.getline(filename, lineno)
else:
name = '[unknown]'
try:
src_lines, src_start = inspect.getsourcelines(frame)
line = src_lines[lineno - src_start]
except (OSError, IndexError):
line = 'Unknown code named [%s]. VM instruction #%d' % (
frame.f_code.co_name, frame.f_lasti)
if self.trace_names is None or name in self.trace_names:
print('%s:%s: %s' % (name, lineno, line.rstrip()))
if not self.show_values:
return self
details = []
tokens = _token_spliter.split(line)
for tok in tokens:
if tok in frame.f_globals:
details.append('%s=%r' % (tok, frame.f_globals[tok]))
if tok in frame.f_locals:
details.append('%s=%r' % (tok, frame.f_locals[tok]))
if details:
print("\t%s" % ' '.join(details))
return self
def spew(trace_names=None, show_values=False):
"""Install a trace hook which writes incredibly detailed logs
about what code is being executed to stdout.
"""
sys.settrace(Spew(trace_names, show_values))
def unspew():
"""Remove the trace hook installed by spew.
"""
sys.settrace(None)
@@ -0,0 +1,81 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Dirty Arbiters - Separate process pool for long-running operations.
Dirty Arbiters provide a separate process pool for executing long-running,
blocking operations (AI model loading, heavy computation) without blocking
HTTP workers. Inspired by Erlang's dirty schedulers.
Key Properties:
- Completely separate from HTTP workers - can be killed/restarted independently
- Stateful - loaded resources persist in dirty worker memory
- Message-passing IPC via Unix sockets with JSON serialization
- Explicit execute() API (no hidden IPC)
- Asyncio-based for clean concurrent handling and future streaming support
"""
from .errors import (
DirtyError,
DirtyTimeoutError,
DirtyConnectionError,
DirtyWorkerError,
DirtyAppError,
DirtyAppNotFoundError,
DirtyProtocolError,
)
from .app import DirtyApp
from .client import (
DirtyClient,
get_dirty_client,
get_dirty_client_async,
set_dirty_socket_path,
close_dirty_client,
close_dirty_client_async,
)
# Stash (shared state between workers)
from . import stash
from .stash import (
StashClient,
StashTable,
StashError,
StashTableNotFoundError,
StashKeyNotFoundError,
)
# Internal imports used by gunicorn core (not part of public API)
from .arbiter import DirtyArbiter
__all__ = [
# Errors
"DirtyError",
"DirtyTimeoutError",
"DirtyConnectionError",
"DirtyWorkerError",
"DirtyAppError",
"DirtyAppNotFoundError",
"DirtyProtocolError",
# App base class
"DirtyApp",
# Client
"DirtyClient",
"get_dirty_client",
"get_dirty_client_async",
"close_dirty_client",
"close_dirty_client_async",
# Stash (shared state)
"stash",
"StashClient",
"StashTable",
"StashError",
"StashTableNotFoundError",
"StashKeyNotFoundError",
# Internal (used by gunicorn core)
"DirtyArbiter",
"set_dirty_socket_path",
]
+350
View File
@@ -0,0 +1,350 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Dirty Application Base Class
Provides the DirtyApp base class that all dirty applications must inherit from,
and utilities for loading dirty apps from import paths.
"""
import importlib
import sys
from .errors import DirtyAppError, DirtyAppNotFoundError
class DirtyApp:
"""
Base class for dirty applications.
Dirty applications are loaded once when the dirty worker starts and
persist in memory for the lifetime of the worker. They are designed
for stateful resources like ML models, connection pools, etc.
Lifecycle
---------
1. ``__init__()``: Called when the app is instantiated (once per worker)
2. ``init()``: Called after instantiation to initialize resources
3. ``__call__()``: Called for each request from HTTP workers
4. ``close()``: Called when the worker shuts down
State Persistence
-----------------
Instance variables persist across requests. This is the key feature
that enables loading heavy resources once and reusing them::
class MLApp(DirtyApp):
def init(self):
self.model = load_model() # Loaded once, reused forever
def predict(self, data):
return self.model.predict(data) # Same model for all requests
Thread Safety
-------------
With ``dirty_threads=1`` (default): Only one request runs at a time,
so no thread safety concerns.
With ``dirty_threads > 1``: Multiple requests may run concurrently
in the same worker. Your app MUST be thread-safe. Options:
- Use locks: ``threading.Lock()`` for shared state
- Use thread-local: ``threading.local()`` for per-thread state
- Use read-only state: Load models once in init(), never mutate
Example::
import threading
class ThreadSafeMLApp(DirtyApp):
def __init__(self):
self.models = {}
self._lock = threading.Lock()
def init(self):
self.models['default'] = load_model('base-model')
def load_model(self, name):
with self._lock:
if name not in self.models:
self.models[name] = load_model(name)
return {"loaded": True, "name": name}
Worker Allocation
-----------------
By default, all dirty workers load all apps. For apps that consume
significant memory (like large ML models), you can limit how many
workers load the app by setting the ``workers`` class attribute::
class HeavyModelApp(DirtyApp):
workers = 2 # Only 2 workers will load this app
def init(self):
self.model = load_10gb_model()
Subclasses should implement:
- init(): Called once at worker startup to initialize resources
- __call__(action, *args, **kwargs): Handle requests from HTTP workers
- close(): Called at worker shutdown to cleanup resources
"""
# Number of workers that should load this app.
# None means all workers (default, backward compatible).
# Set to an integer to limit how many workers load this app.
workers = None
def init(self):
"""
Initialize the application.
Called once when the dirty worker starts, after the app instance
is created. Use this for expensive initialization like loading
ML models, establishing database connections, etc.
This method is called in the child process after fork, so it's
safe to initialize non-fork-safe resources here.
"""
def __call__(self, action, *args, **kwargs):
"""
Handle a request from an HTTP worker.
Args:
action: The action/method name to execute
*args: Positional arguments for the action
**kwargs: Keyword arguments for the action
Returns:
The result of the action (must be JSON-serializable)
Raises:
ValueError: If the action is unknown
Any exception: Will be caught and returned as DirtyAppError
"""
method = getattr(self, action, None)
if method is None or action.startswith('_'):
raise ValueError(f"Unknown action: {action}")
return method(*args, **kwargs)
def close(self):
"""
Cleanup resources.
Called when the dirty worker is shutting down. Use this to
release resources like database connections, unload models, etc.
"""
def parse_dirty_app_spec(spec):
"""
Parse a dirty app specification.
Supports two formats:
- ``"module:Class"`` - standard format, all workers load the app
- ``"module:Class:N"`` - worker-limited format, only N workers load the app
Args:
spec: The app specification string
Returns:
tuple: (import_path, worker_count)
- import_path: The "module:Class" part for importing
- worker_count: Integer limit or None for all workers
Raises:
DirtyAppError: If the spec format is invalid or worker_count is < 1
Examples::
>>> parse_dirty_app_spec("myapp:App")
("myapp:App", None)
>>> parse_dirty_app_spec("myapp:App:2")
("myapp:App", 2)
>>> parse_dirty_app_spec("myapp.sub:App:1")
("myapp.sub:App", 1)
"""
if ':' not in spec:
raise DirtyAppError(
f"Invalid import path format: {spec}. "
f"Expected 'module.path:ClassName' or 'module.path:ClassName:N'",
app_path=spec
)
parts = spec.split(':')
# Standard format: "module:Class" or "module.sub:Class"
if len(parts) == 2:
return (spec, None)
# Worker-limited format: "module:Class:N"
if len(parts) == 3:
module_path, class_name, count_str = parts
import_path = f"{module_path}:{class_name}"
# Validate the worker count
try:
worker_count = int(count_str)
except ValueError:
raise DirtyAppError(
f"Invalid worker count in spec: {spec}. "
f"Expected integer, got '{count_str}'",
app_path=spec
)
if worker_count < 1:
raise DirtyAppError(
f"Invalid worker count in spec: {spec}. "
f"Worker count must be >= 1, got {worker_count}",
app_path=spec
)
return (import_path, worker_count)
# Too many colons
raise DirtyAppError(
f"Invalid import path format: {spec}. "
f"Expected 'module.path:ClassName' or 'module.path:ClassName:N'",
app_path=spec
)
def load_dirty_app(import_path):
"""
Load a dirty app class from an import path.
Args:
import_path: String in format 'module.path:ClassName'
Returns:
An instance of the dirty app class
Raises:
DirtyAppNotFoundError: If the module or class cannot be found
DirtyAppError: If the class is not a valid DirtyApp subclass
"""
if ':' not in import_path:
raise DirtyAppError(
f"Invalid import path format: {import_path}. "
f"Expected 'module.path:ClassName'",
app_path=import_path
)
module_path, class_name = import_path.rsplit(':', 1)
try:
# Import the module
if module_path in sys.modules:
module = sys.modules[module_path]
else:
module = importlib.import_module(module_path)
except ImportError as e:
raise DirtyAppNotFoundError(import_path) from e
# Get the class from the module
try:
app_class = getattr(module, class_name)
except AttributeError:
raise DirtyAppNotFoundError(import_path) from None
# Validate it's a class
if not isinstance(app_class, type):
raise DirtyAppError(
f"{import_path} is not a class",
app_path=import_path
)
# Create an instance
try:
app = app_class()
except Exception as e:
raise DirtyAppError(
f"Failed to instantiate {import_path}: {e}",
app_path=import_path
) from e
# Validate it has the required methods
required_methods = ['init', '__call__', 'close']
for method_name in required_methods:
if not hasattr(app, method_name) or not callable(getattr(app, method_name)):
raise DirtyAppError(
f"{import_path} is missing required method: {method_name}",
app_path=import_path
)
return app
def load_dirty_apps(import_paths):
"""
Load multiple dirty apps from a list of import paths.
Args:
import_paths: List of import path strings
Returns:
dict: Mapping of import path to app instance
Raises:
DirtyAppError: If any app fails to load
"""
apps = {}
for import_path in import_paths:
apps[import_path] = load_dirty_app(import_path)
return apps
def get_app_workers_attribute(import_path):
"""
Get the workers class attribute from a dirty app without instantiating it.
This is used by the arbiter to determine how many workers should load
an app based on the class attribute, without needing to actually load
the app.
Args:
import_path: String in format 'module.path:ClassName'
Returns:
The workers class attribute value (int or None)
Raises:
DirtyAppNotFoundError: If the module or class cannot be found
DirtyAppError: If the import path format is invalid
"""
if ':' not in import_path:
raise DirtyAppError(
f"Invalid import path format: {import_path}. "
f"Expected 'module.path:ClassName'",
app_path=import_path
)
module_path, class_name = import_path.rsplit(':', 1)
try:
# Import the module
if module_path in sys.modules:
module = sys.modules[module_path]
else:
module = importlib.import_module(module_path)
except ImportError as e:
raise DirtyAppNotFoundError(import_path) from e
# Get the class from the module
try:
app_class = getattr(module, class_name)
except AttributeError:
raise DirtyAppNotFoundError(import_path) from None
# Validate it's a class
if not isinstance(app_class, type):
raise DirtyAppError(
f"{import_path} is not a class",
app_path=import_path
)
# Return the workers attribute (defaults to None if not set)
return getattr(app_class, 'workers', None)
File diff suppressed because it is too large Load Diff
+754
View File
@@ -0,0 +1,754 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Dirty Client
Client for HTTP workers to communicate with the dirty worker pool.
Provides both sync and async APIs.
"""
import asyncio
import contextvars
import os
import socket
import threading
import time
import uuid
from .errors import (
DirtyConnectionError,
DirtyError,
DirtyTimeoutError,
)
from .protocol import (
DirtyProtocol,
make_request,
)
class DirtyClient:
"""
Client for calling dirty workers from HTTP workers.
Provides both sync and async APIs. The sync API is for traditional
sync workers (sync, gthread), while the async API is for async
workers (asgi, gevent).
"""
def __init__(self, socket_path, timeout=30.0):
"""
Initialize the dirty client.
Args:
socket_path: Path to the dirty arbiter's Unix socket
timeout: Default timeout for operations in seconds
"""
self.socket_path = socket_path
self.timeout = timeout
self._sock = None
self._reader = None
self._writer = None
self._lock = threading.Lock()
# -------------------------------------------------------------------------
# Sync API (for sync HTTP workers)
# -------------------------------------------------------------------------
def connect(self):
"""
Establish sync socket connection to arbiter.
Raises:
DirtyConnectionError: If connection fails
"""
if self._sock is not None:
return
try:
self._sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
self._sock.settimeout(self.timeout)
self._sock.connect(self.socket_path)
except (socket.error, OSError) as e:
self._sock = None
raise DirtyConnectionError(
f"Failed to connect to dirty arbiter: {e}",
socket_path=self.socket_path
) from e
def execute(self, app_path, action, *args, **kwargs):
"""
Execute an action on a dirty app (sync/blocking).
Args:
app_path: Import path of the dirty app (e.g., 'myapp.ml:MLApp')
action: Action to call on the app
*args: Positional arguments
**kwargs: Keyword arguments
Returns:
Result from the dirty app action
Raises:
DirtyConnectionError: If connection fails
DirtyTimeoutError: If operation times out
DirtyError: If execution fails
"""
with self._lock:
return self._execute_locked(app_path, action, args, kwargs)
def _execute_locked(self, app_path, action, args, kwargs):
"""Execute while holding the lock."""
# Ensure connected
if self._sock is None:
self.connect()
# Build request
request_id = str(uuid.uuid4())
request = make_request(
request_id=request_id,
app_path=app_path,
action=action,
args=args,
kwargs=kwargs
)
try:
# Send request
DirtyProtocol.write_message(self._sock, request)
# Receive response
response = DirtyProtocol.read_message(self._sock)
# Handle response
return self._handle_response(response)
except socket.timeout:
self._close_socket()
raise DirtyTimeoutError(
"Timeout waiting for dirty app response",
timeout=self.timeout
)
except Exception as e:
self._close_socket()
if isinstance(e, DirtyError):
raise
raise DirtyConnectionError(f"Communication error: {e}") from e
def stream(self, app_path, action, *args, **kwargs):
"""
Stream results from a dirty app action (sync).
This method returns an iterator that yields chunks from a streaming
response. Use this for actions that return generators.
Args:
app_path: Import path of the dirty app (e.g., 'myapp.ml:MLApp')
action: Action to call on the app
*args: Positional arguments
**kwargs: Keyword arguments
Yields:
Chunks of data from the streaming response
Raises:
DirtyConnectionError: If connection fails
DirtyTimeoutError: If operation times out
DirtyError: If execution fails
Example::
for chunk in client.stream("myapp.llm:LLMApp", "generate", prompt):
print(chunk, end="", flush=True)
"""
return DirtyStreamIterator(self, app_path, action, args, kwargs)
def _handle_response(self, response):
"""Handle response message, extracting result or raising error."""
msg_type = response.get("type")
if msg_type == DirtyProtocol.MSG_TYPE_RESPONSE:
return response.get("result")
elif msg_type == DirtyProtocol.MSG_TYPE_ERROR:
error_info = response.get("error", {})
error = DirtyError.from_dict(error_info)
raise error
else:
raise DirtyError(f"Unknown response type: {msg_type}")
def _close_socket(self):
"""Close the socket connection."""
if self._sock is not None:
try:
self._sock.close()
except Exception:
pass
self._sock = None
def close(self):
"""Close the sync connection."""
with self._lock:
self._close_socket()
# -------------------------------------------------------------------------
# Async API (for async HTTP workers)
# -------------------------------------------------------------------------
async def connect_async(self):
"""
Establish async connection to arbiter.
Raises:
DirtyConnectionError: If connection fails
"""
if self._writer is not None:
return
try:
self._reader, self._writer = await asyncio.wait_for(
asyncio.open_unix_connection(self.socket_path),
timeout=self.timeout
)
except asyncio.TimeoutError:
raise DirtyTimeoutError(
"Timeout connecting to dirty arbiter",
timeout=self.timeout
)
except (OSError, ConnectionError) as e:
raise DirtyConnectionError(
f"Failed to connect to dirty arbiter: {e}",
socket_path=self.socket_path
) from e
async def execute_async(self, app_path, action, *args, **kwargs):
"""
Execute an action on a dirty app (async/non-blocking).
Args:
app_path: Import path of the dirty app
action: Action to call on the app
*args: Positional arguments
**kwargs: Keyword arguments
Returns:
Result from the dirty app action
Raises:
DirtyConnectionError: If connection fails
DirtyTimeoutError: If operation times out
DirtyError: If execution fails
"""
# Ensure connected
if self._writer is None:
await self.connect_async()
# Build request
request_id = str(uuid.uuid4())
request = make_request(
request_id=request_id,
app_path=app_path,
action=action,
args=args,
kwargs=kwargs
)
try:
# Send request
await DirtyProtocol.write_message_async(self._writer, request)
# Receive response with timeout
response = await asyncio.wait_for(
DirtyProtocol.read_message_async(self._reader),
timeout=self.timeout
)
# Handle response
return self._handle_response(response)
except asyncio.TimeoutError:
await self._close_async()
raise DirtyTimeoutError(
"Timeout waiting for dirty app response",
timeout=self.timeout
)
except Exception as e:
await self._close_async()
if isinstance(e, DirtyError):
raise
raise DirtyConnectionError(f"Communication error: {e}") from e
def stream_async(self, app_path, action, *args, **kwargs):
"""
Stream results from a dirty app action (async).
This method returns an async iterator that yields chunks from a
streaming response. Use this for actions that return generators.
Args:
app_path: Import path of the dirty app (e.g., 'myapp.ml:MLApp')
action: Action to call on the app
*args: Positional arguments
**kwargs: Keyword arguments
Yields:
Chunks of data from the streaming response
Raises:
DirtyConnectionError: If connection fails
DirtyTimeoutError: If operation times out
DirtyError: If execution fails
Example::
async for chunk in client.stream_async("myapp.llm:LLMApp", "generate", prompt):
await response.write(chunk)
"""
return DirtyAsyncStreamIterator(self, app_path, action, args, kwargs)
async def _close_async(self):
"""Close the async connection."""
if self._writer is not None:
try:
self._writer.close()
await self._writer.wait_closed()
except Exception:
pass
self._writer = None
self._reader = None
async def close_async(self):
"""Close the async connection."""
await self._close_async()
# -------------------------------------------------------------------------
# Context managers
# -------------------------------------------------------------------------
def __enter__(self):
self.connect()
return self
def __exit__(self, exc_type, exc_val, exc_tb):
self.close()
async def __aenter__(self):
await self.connect_async()
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
await self.close_async()
# =============================================================================
# Stream Iterator classes
# =============================================================================
class DirtyStreamIterator:
"""
Iterator for streaming responses from dirty workers (sync).
This class is returned by `DirtyClient.stream()` and yields chunks
from a streaming response until the end message is received.
Uses a deadline-based timeout approach:
- Total stream timeout: limits entire stream duration
- Idle timeout: limits gap between chunks (defaults to total timeout)
"""
# Default idle timeout between chunks (seconds)
DEFAULT_IDLE_TIMEOUT = 30.0
# Threshold for applying per-read timeout (seconds)
# When remaining time is above this, use a larger timeout for efficiency
_TIMEOUT_THRESHOLD = 5.0
def __init__(self, client, app_path, action, args, kwargs,
idle_timeout=None):
self.client = client
self.app_path = app_path
self.action = action
self.args = args
self.kwargs = kwargs
self._started = False
self._exhausted = False
self._request_id = None
self._deadline = None
self._last_chunk_time = None
# Idle timeout: max time between chunks
self._idle_timeout = (
idle_timeout if idle_timeout is not None
else min(self.DEFAULT_IDLE_TIMEOUT, client.timeout)
)
def __iter__(self):
return self
def __next__(self):
if self._exhausted:
raise StopIteration
if not self._started:
self._start_request()
self._started = True
return self._read_next_chunk()
def _start_request(self):
"""Send the initial request to the arbiter."""
with self.client._lock:
if self.client._sock is None:
self.client.connect()
# Set deadline for entire stream
now = time.monotonic()
self._deadline = now + self.client.timeout
self._last_chunk_time = now
self._request_id = str(uuid.uuid4())
request = make_request(
self._request_id,
self.app_path,
self.action,
args=self.args,
kwargs=self.kwargs,
)
DirtyProtocol.write_message(self.client._sock, request)
def _read_next_chunk(self):
"""Read the next message from the stream."""
with self.client._lock:
# Check total stream deadline
now = time.monotonic()
if now >= self._deadline:
self._exhausted = True
raise DirtyTimeoutError(
"Stream exceeded total timeout",
timeout=self.client.timeout
)
remaining = self._deadline - now
# Set socket timeout based on remaining time
# Fast path: use larger timeout when plenty of time remains
if remaining > self._TIMEOUT_THRESHOLD:
read_timeout = self._TIMEOUT_THRESHOLD
else:
read_timeout = min(remaining, self._idle_timeout)
try:
self.client._sock.settimeout(read_timeout)
response = DirtyProtocol.read_message(self.client._sock)
except socket.timeout:
# Check which timeout was hit
now = time.monotonic()
if now >= self._deadline:
self._exhausted = True
raise DirtyTimeoutError(
"Stream exceeded total timeout",
timeout=self.client.timeout
)
idle_duration = now - self._last_chunk_time
self._exhausted = True
raise DirtyTimeoutError(
f"Timeout waiting for next chunk (idle {idle_duration:.1f}s)",
timeout=self._idle_timeout
)
except Exception as e:
self._exhausted = True
self.client._close_socket()
raise DirtyConnectionError(f"Communication error: {e}") from e
# Update last chunk time for idle tracking
self._last_chunk_time = time.monotonic()
msg_type = response.get("type")
# Chunk message - return the data
if msg_type == DirtyProtocol.MSG_TYPE_CHUNK:
return response.get("data")
# End message - stop iteration
if msg_type == DirtyProtocol.MSG_TYPE_END:
self._exhausted = True
raise StopIteration
# Error message - raise exception
if msg_type == DirtyProtocol.MSG_TYPE_ERROR:
self._exhausted = True
error_info = response.get("error", {})
raise DirtyError.from_dict(error_info)
# Regular response - shouldn't happen for streaming, but handle it
if msg_type == DirtyProtocol.MSG_TYPE_RESPONSE:
self._exhausted = True
# Return the result as the only chunk then stop
raise StopIteration
# Unknown type
self._exhausted = True
raise DirtyError(f"Unknown message type: {msg_type}")
class DirtyAsyncStreamIterator:
"""
Async iterator for streaming responses from dirty workers.
This class is returned by `DirtyClient.stream_async()` and yields chunks
from a streaming response until the end message is received.
Uses a deadline-based timeout approach for efficiency:
- Total stream timeout: limits entire stream duration
- Idle timeout: limits gap between chunks (defaults to total timeout)
This avoids the overhead of asyncio.wait_for() on every chunk read.
"""
# Default idle timeout between chunks (seconds)
DEFAULT_IDLE_TIMEOUT = 30.0
def __init__(self, client, app_path, action, args, kwargs,
idle_timeout=None):
self.client = client
self.app_path = app_path
self.action = action
self.args = args
self.kwargs = kwargs
self._started = False
self._exhausted = False
self._request_id = None
self._deadline = None
self._last_chunk_time = None
# Idle timeout: max time between chunks
self._idle_timeout = (
idle_timeout if idle_timeout is not None
else min(self.DEFAULT_IDLE_TIMEOUT, client.timeout)
)
def __aiter__(self):
return self
async def __anext__(self):
if self._exhausted:
raise StopAsyncIteration
if not self._started:
await self._start_request()
self._started = True
return await self._read_next_chunk()
async def _start_request(self):
"""Send the initial request to the arbiter."""
if self.client._writer is None:
await self.client.connect_async()
# Set deadline for entire stream
now = time.monotonic()
self._deadline = now + self.client.timeout
self._last_chunk_time = now
self._request_id = str(uuid.uuid4())
request = make_request(
self._request_id,
self.app_path,
self.action,
args=self.args,
kwargs=self.kwargs,
)
await DirtyProtocol.write_message_async(self.client._writer, request)
# Threshold for applying timeout wrapper (seconds)
# When remaining time is above this, skip timeout for performance
_TIMEOUT_THRESHOLD = 5.0
async def _read_next_chunk(self):
"""Read the next message from the stream."""
# Calculate remaining time until deadline
now = time.monotonic()
# Check total stream deadline
if now >= self._deadline:
self._exhausted = True
raise DirtyTimeoutError(
"Stream exceeded total timeout",
timeout=self.client.timeout
)
remaining = self._deadline - now
try:
# Fast path: skip timeout wrapper when we have plenty of time
# This avoids asyncio.wait_for() overhead for most chunks
if remaining > self._TIMEOUT_THRESHOLD:
response = await DirtyProtocol.read_message_async(
self.client._reader
)
else:
# Near deadline: apply timeout protection
read_timeout = min(remaining, self._idle_timeout)
response = await asyncio.wait_for(
DirtyProtocol.read_message_async(self.client._reader),
timeout=read_timeout
)
except asyncio.TimeoutError:
self._exhausted = True
now = time.monotonic()
if now >= self._deadline:
raise DirtyTimeoutError(
"Stream exceeded total timeout",
timeout=self.client.timeout
)
idle_duration = now - self._last_chunk_time
raise DirtyTimeoutError(
f"Timeout waiting for next chunk (idle {idle_duration:.1f}s)",
timeout=self._idle_timeout
)
except Exception as e:
self._exhausted = True
await self.client._close_async()
raise DirtyConnectionError(f"Communication error: {e}") from e
# Update last chunk time for idle tracking
self._last_chunk_time = time.monotonic()
msg_type = response.get("type")
# Chunk message - return the data
if msg_type == DirtyProtocol.MSG_TYPE_CHUNK:
return response.get("data")
# End message - stop iteration
if msg_type == DirtyProtocol.MSG_TYPE_END:
self._exhausted = True
raise StopAsyncIteration
# Error message - raise exception
if msg_type == DirtyProtocol.MSG_TYPE_ERROR:
self._exhausted = True
error_info = response.get("error", {})
raise DirtyError.from_dict(error_info)
# Regular response - shouldn't happen for streaming
if msg_type == DirtyProtocol.MSG_TYPE_RESPONSE:
self._exhausted = True
raise StopAsyncIteration
# Unknown type
self._exhausted = True
raise DirtyError(f"Unknown message type: {msg_type}")
# =============================================================================
# Thread-local and context-local client management
# =============================================================================
# Thread-local storage for sync workers
_thread_local = threading.local()
# Context var for async workers
_async_client_var: contextvars.ContextVar[DirtyClient] = contextvars.ContextVar(
'dirty_client'
)
# Global socket path (set by arbiter)
_dirty_socket_path = None
def set_dirty_socket_path(path):
"""Set the global dirty socket path (called during initialization)."""
global _dirty_socket_path # pylint: disable=global-statement
_dirty_socket_path = path
# Also set the stash socket path (uses same arbiter socket)
from .stash import set_stash_socket_path
set_stash_socket_path(path)
def get_dirty_socket_path():
"""Get the dirty socket path."""
if _dirty_socket_path is None:
# Check environment variable
path = os.environ.get('GUNICORN_DIRTY_SOCKET')
if path:
return path
raise DirtyError(
"Dirty socket path not configured. "
"Make sure dirty_workers > 0 and dirty_apps are configured."
)
return _dirty_socket_path
def get_dirty_client(timeout=30.0) -> DirtyClient:
"""
Get or create a thread-local sync client.
This is the recommended way to get a client in sync HTTP workers.
Args:
timeout: Timeout for operations in seconds
Returns:
DirtyClient: Thread-local client instance
Example::
from gunicorn.dirty import get_dirty_client
def my_view(request):
client = get_dirty_client()
result = client.execute("myapp.ml:MLApp", "inference", data)
return result
"""
client = getattr(_thread_local, 'dirty_client', None)
if client is None:
socket_path = get_dirty_socket_path()
client = DirtyClient(socket_path, timeout=timeout)
_thread_local.dirty_client = client
return client
async def get_dirty_client_async(timeout=30.0) -> DirtyClient:
"""
Get or create a context-local async client.
This is the recommended way to get a client in async HTTP workers.
Args:
timeout: Timeout for operations in seconds
Returns:
DirtyClient: Context-local client instance
Example::
from gunicorn.dirty import get_dirty_client_async
async def my_view(request):
client = await get_dirty_client_async()
result = await client.execute_async("myapp.ml:MLApp", "inference", data)
return result
"""
try:
client = _async_client_var.get()
except LookupError:
socket_path = get_dirty_socket_path()
client = DirtyClient(socket_path, timeout=timeout)
_async_client_var.set(client)
return client
def close_dirty_client():
"""Close the thread-local client (call on worker exit)."""
client = getattr(_thread_local, 'dirty_client', None)
if client is not None:
client.close()
_thread_local.dirty_client = None
async def close_dirty_client_async():
"""Close the context-local async client."""
try:
client = _async_client_var.get()
await client.close_async()
except LookupError:
pass
+180
View File
@@ -0,0 +1,180 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Dirty Arbiters Error Classes
Exception hierarchy for dirty worker pool operations.
"""
class DirtyError(Exception):
"""Base exception for all dirty arbiter errors."""
def __init__(self, message, details=None):
self.message = message
self.details = details or {}
super().__init__(message)
def __str__(self):
if self.details:
return f"{self.message}: {self.details}"
return self.message
def to_dict(self):
"""Serialize error for protocol transmission."""
return {
"error_type": self.__class__.__name__,
"message": self.message,
"details": self.details,
}
@classmethod
def from_dict(cls, data):
"""Deserialize error from protocol transmission.
Creates an error instance from a serialized dict. The returned
error will be an instance of the appropriate subclass based on
the error_type field, but constructed using the base DirtyError
__init__ to preserve all details.
"""
error_classes = {
"DirtyError": DirtyError,
"DirtyTimeoutError": DirtyTimeoutError,
"DirtyConnectionError": DirtyConnectionError,
"DirtyWorkerError": DirtyWorkerError,
"DirtyAppError": DirtyAppError,
"DirtyAppNotFoundError": DirtyAppNotFoundError,
"DirtyNoWorkersAvailableError": DirtyNoWorkersAvailableError,
"DirtyProtocolError": DirtyProtocolError,
}
error_type = data.get("error_type", "DirtyError")
error_class = error_classes.get(error_type, DirtyError)
# Create instance and set attributes directly to bypass
# subclass __init__ complexity while preserving error type
error = Exception.__new__(error_class)
error.message = data.get("message", "Unknown error")
error.details = data.get("details") or {}
Exception.__init__(error, error.message)
# Set subclass-specific attributes from details
if error_class == DirtyTimeoutError:
error.timeout = error.details.get("timeout")
elif error_class == DirtyConnectionError:
error.socket_path = error.details.get("socket_path")
elif error_class == DirtyWorkerError:
error.worker_id = error.details.get("worker_id")
error.traceback = error.details.get("traceback")
elif error_class in (DirtyAppError, DirtyAppNotFoundError):
error.app_path = error.details.get("app_path")
error.action = error.details.get("action")
error.traceback = error.details.get("traceback")
elif error_class == DirtyNoWorkersAvailableError:
error.app_path = error.details.get("app_path")
return error
class DirtyTimeoutError(DirtyError):
"""Raised when a dirty operation times out."""
def __init__(self, message="Operation timed out", timeout=None):
details = {"timeout": timeout} if timeout else {}
super().__init__(message, details)
self.timeout = timeout
class DirtyConnectionError(DirtyError):
"""Raised when connection to dirty arbiter fails."""
def __init__(self, message="Connection failed", socket_path=None):
details = {"socket_path": socket_path} if socket_path else {}
super().__init__(message, details)
self.socket_path = socket_path
class DirtyWorkerError(DirtyError):
"""Raised when a dirty worker encounters an error."""
def __init__(self, message, worker_id=None, traceback=None):
details = {}
if worker_id is not None:
details["worker_id"] = worker_id
if traceback:
details["traceback"] = traceback
super().__init__(message, details)
self.worker_id = worker_id
self.traceback = traceback
class DirtyAppError(DirtyError):
"""Raised when a dirty app encounters an error during execution."""
def __init__(self, message, app_path=None, action=None, traceback=None):
details = {}
if app_path:
details["app_path"] = app_path
if action:
details["action"] = action
if traceback:
details["traceback"] = traceback
super().__init__(message, details)
self.app_path = app_path
self.action = action
self.traceback = traceback
class DirtyAppNotFoundError(DirtyAppError):
"""Raised when a dirty app is not found."""
def __init__(self, app_path):
super().__init__(f"Dirty app not found: {app_path}", app_path=app_path)
class DirtyNoWorkersAvailableError(DirtyError):
"""
Raised when no workers are available for the requested app.
This exception is raised when a request targets an app that has
worker limits configured, and no workers with that app are currently
available (e.g., all workers for that app crashed and haven't been
respawned yet).
Web applications can catch this exception to provide graceful
degradation, such as queuing requests for retry or showing a
maintenance page.
Example::
from gunicorn.dirty import get_dirty_client
from gunicorn.dirty.errors import DirtyNoWorkersAvailableError
def my_view(request):
client = get_dirty_client()
try:
result = client.execute("myapp.ml:HeavyModel", "predict", data)
except DirtyNoWorkersAvailableError as e:
return {"error": "Service temporarily unavailable",
"app": e.app_path}
"""
def __init__(self, app_path, message=None):
if message is None:
message = f"No workers available for app: {app_path}"
super().__init__(message, details={"app_path": app_path})
self.app_path = app_path
class DirtyProtocolError(DirtyError):
"""Raised when there is a protocol-level error."""
def __init__(self, message="Protocol error", raw_data=None):
details = {}
if raw_data is not None:
# Truncate raw data for safety
if isinstance(raw_data, bytes):
raw_data = raw_data[:100].hex()
details["raw_data"] = str(raw_data)[:200]
super().__init__(message, details)
@@ -0,0 +1,810 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Dirty Worker Binary Protocol
Binary message framing over Unix sockets, inspired by OpenBSD msgctl/msgsnd.
Replaces JSON protocol for efficient binary data transfer.
Header Format (16 bytes):
+--------+--------+--------+--------+--------+--------+--------+--------+
| Magic (2B) | Ver(1) | MType | Payload Length (4B) |
+--------+--------+--------+--------+--------+--------+--------+--------+
| Request ID (8 bytes) |
+--------+--------+--------+--------+--------+--------+--------+--------+
- Magic: 0x47 0x44 ("GD" for Gunicorn Dirty)
- Version: 0x01
- MType: Message type (REQUEST, RESPONSE, ERROR, CHUNK, END)
- Length: Payload size (big-endian uint32, max 64MB)
- Request ID: uint64 (replaces UUID string)
Payload is TLV-encoded (see tlv.py).
"""
import asyncio
import socket
import struct
from .errors import DirtyProtocolError
from .tlv import TLVEncoder
# Protocol constants
MAGIC = b"GD" # 0x47 0x44
VERSION = 0x01
# Message types (1 byte)
MSG_TYPE_REQUEST = 0x01
MSG_TYPE_RESPONSE = 0x02
MSG_TYPE_ERROR = 0x03
MSG_TYPE_CHUNK = 0x04
MSG_TYPE_END = 0x05
MSG_TYPE_STASH = 0x10 # Stash operations (shared state between workers)
MSG_TYPE_STATUS = 0x11 # Status query for arbiter/workers
MSG_TYPE_MANAGE = 0x12 # Worker management (add/remove workers)
# Message type names (for backwards compatibility with old API)
MSG_TYPE_REQUEST_STR = "request"
MSG_TYPE_RESPONSE_STR = "response"
MSG_TYPE_ERROR_STR = "error"
MSG_TYPE_CHUNK_STR = "chunk"
MSG_TYPE_END_STR = "end"
MSG_TYPE_STASH_STR = "stash"
MSG_TYPE_STATUS_STR = "status"
MSG_TYPE_MANAGE_STR = "manage"
# Map int types to string names
MSG_TYPE_TO_STR = {
MSG_TYPE_REQUEST: MSG_TYPE_REQUEST_STR,
MSG_TYPE_RESPONSE: MSG_TYPE_RESPONSE_STR,
MSG_TYPE_ERROR: MSG_TYPE_ERROR_STR,
MSG_TYPE_CHUNK: MSG_TYPE_CHUNK_STR,
MSG_TYPE_END: MSG_TYPE_END_STR,
MSG_TYPE_STASH: MSG_TYPE_STASH_STR,
MSG_TYPE_STATUS: MSG_TYPE_STATUS_STR,
MSG_TYPE_MANAGE: MSG_TYPE_MANAGE_STR,
}
# Map string names to int types
MSG_TYPE_FROM_STR = {v: k for k, v in MSG_TYPE_TO_STR.items()}
# Stash operation codes
STASH_OP_PUT = 1
STASH_OP_GET = 2
STASH_OP_DELETE = 3
STASH_OP_KEYS = 4
STASH_OP_CLEAR = 5
STASH_OP_INFO = 6
STASH_OP_ENSURE = 7
STASH_OP_DELETE_TABLE = 8
STASH_OP_TABLES = 9
STASH_OP_EXISTS = 10
# Manage operation codes
MANAGE_OP_ADD = 1 # Add/spawn workers
MANAGE_OP_REMOVE = 2 # Remove/kill workers
# Header format: Magic (2) + Version (1) + Type (1) + Length (4) + RequestID (8) = 16
HEADER_FORMAT = ">2sBBIQ"
HEADER_SIZE = struct.calcsize(HEADER_FORMAT)
# Maximum message size (64 MB)
MAX_MESSAGE_SIZE = 64 * 1024 * 1024
class BinaryProtocol:
"""Binary message protocol for dirty worker IPC."""
# Export constants for external use
HEADER_SIZE = HEADER_SIZE
MAX_MESSAGE_SIZE = MAX_MESSAGE_SIZE
MSG_TYPE_REQUEST = MSG_TYPE_REQUEST_STR
MSG_TYPE_RESPONSE = MSG_TYPE_RESPONSE_STR
MSG_TYPE_ERROR = MSG_TYPE_ERROR_STR
MSG_TYPE_CHUNK = MSG_TYPE_CHUNK_STR
MSG_TYPE_END = MSG_TYPE_END_STR
MSG_TYPE_STASH = MSG_TYPE_STASH_STR
MSG_TYPE_STATUS = MSG_TYPE_STATUS_STR
MSG_TYPE_MANAGE = MSG_TYPE_MANAGE_STR
@staticmethod
def encode_header(msg_type: int, request_id: int, payload_length: int) -> bytes:
"""
Encode the 16-byte message header.
Args:
msg_type: Message type (MSG_TYPE_REQUEST, etc.)
request_id: Unique request identifier (uint64)
payload_length: Length of the TLV-encoded payload
Returns:
bytes: 16-byte header
"""
return struct.pack(HEADER_FORMAT, MAGIC, VERSION, msg_type,
payload_length, request_id)
@staticmethod
def decode_header(data: bytes) -> tuple:
"""
Decode the 16-byte message header.
Args:
data: 16 bytes of header data
Returns:
tuple: (msg_type, request_id, payload_length)
Raises:
DirtyProtocolError: If header is invalid
"""
if len(data) < HEADER_SIZE:
raise DirtyProtocolError(
f"Header too short: {len(data)} bytes, expected {HEADER_SIZE}",
raw_data=data
)
magic, version, msg_type, length, request_id = struct.unpack(
HEADER_FORMAT, data[:HEADER_SIZE]
)
if magic != MAGIC:
raise DirtyProtocolError(
f"Invalid magic: {magic!r}, expected {MAGIC!r}",
raw_data=data[:20]
)
if version != VERSION:
raise DirtyProtocolError(
f"Unsupported protocol version: {version}, expected {VERSION}",
raw_data=data[:20]
)
if msg_type not in MSG_TYPE_TO_STR:
raise DirtyProtocolError(
f"Unknown message type: 0x{msg_type:02x}",
raw_data=data[:20]
)
if length > MAX_MESSAGE_SIZE:
raise DirtyProtocolError(
f"Message too large: {length} bytes (max: {MAX_MESSAGE_SIZE})"
)
return msg_type, request_id, length
@staticmethod
def encode_request(request_id: int, app_path: str, action: str,
args: tuple = None, kwargs: dict = None) -> bytes:
"""
Encode a request message.
Args:
request_id: Unique request identifier (uint64)
app_path: Import path of the dirty app
action: Action to call on the app
args: Positional arguments
kwargs: Keyword arguments
Returns:
bytes: Complete message (header + payload)
"""
payload_dict = {
"app_path": app_path,
"action": action,
"args": list(args) if args else [],
"kwargs": kwargs or {},
}
payload = TLVEncoder.encode(payload_dict)
header = BinaryProtocol.encode_header(MSG_TYPE_REQUEST, request_id,
len(payload))
return header + payload
@staticmethod
def encode_response(request_id: int, result) -> bytes:
"""
Encode a success response message.
Args:
request_id: Request identifier this responds to
result: Result value (must be TLV-serializable)
Returns:
bytes: Complete message (header + payload)
"""
payload_dict = {"result": result}
payload = TLVEncoder.encode(payload_dict)
header = BinaryProtocol.encode_header(MSG_TYPE_RESPONSE, request_id,
len(payload))
return header + payload
@staticmethod
def encode_error(request_id: int, error) -> bytes:
"""
Encode an error response message.
Args:
request_id: Request identifier this responds to
error: DirtyError instance, dict, or Exception
Returns:
bytes: Complete message (header + payload)
"""
from .errors import DirtyError
if isinstance(error, DirtyError):
error_dict = error.to_dict()
elif isinstance(error, dict):
error_dict = error
else:
error_dict = {
"error_type": type(error).__name__,
"message": str(error),
"details": {},
}
payload_dict = {"error": error_dict}
payload = TLVEncoder.encode(payload_dict)
header = BinaryProtocol.encode_header(MSG_TYPE_ERROR, request_id,
len(payload))
return header + payload
@staticmethod
def encode_chunk(request_id: int, data) -> bytes:
"""
Encode a chunk message for streaming responses.
Args:
request_id: Request identifier this chunk belongs to
data: Chunk data (must be TLV-serializable)
Returns:
bytes: Complete message (header + payload)
"""
payload_dict = {"data": data}
payload = TLVEncoder.encode(payload_dict)
header = BinaryProtocol.encode_header(MSG_TYPE_CHUNK, request_id,
len(payload))
return header + payload
@staticmethod
def encode_end(request_id: int) -> bytes:
"""
Encode an end-of-stream message.
Args:
request_id: Request identifier this ends
Returns:
bytes: Complete message (header + empty payload)
"""
# End message has empty payload
header = BinaryProtocol.encode_header(MSG_TYPE_END, request_id, 0)
return header
@staticmethod
def encode_status(request_id: int) -> bytes:
"""
Encode a status query message.
Args:
request_id: Request identifier
Returns:
bytes: Complete message (header + empty payload)
"""
# Status query has empty payload
header = BinaryProtocol.encode_header(MSG_TYPE_STATUS, request_id, 0)
return header
@staticmethod
def encode_manage(request_id: int, op: int, count: int = 1) -> bytes:
"""
Encode a worker management message.
Args:
request_id: Request identifier
op: Management operation (MANAGE_OP_ADD or MANAGE_OP_REMOVE)
count: Number of workers to add/remove
Returns:
bytes: Complete message (header + payload)
"""
payload_dict = {
"op": op,
"count": count,
}
payload = TLVEncoder.encode(payload_dict)
header = BinaryProtocol.encode_header(MSG_TYPE_MANAGE, request_id,
len(payload))
return header + payload
@staticmethod
def encode_stash(request_id: int, op: int, table: str,
key=None, value=None, pattern=None) -> bytes:
"""
Encode a stash operation message.
Args:
request_id: Unique request identifier (uint64)
op: Stash operation code (STASH_OP_*)
table: Table name
key: Optional key for put/get/delete operations
value: Optional value for put operation
pattern: Optional pattern for keys operation
Returns:
bytes: Complete message (header + payload)
"""
payload_dict = {
"op": op,
"table": table,
}
if key is not None:
payload_dict["key"] = key
if value is not None:
payload_dict["value"] = value
if pattern is not None:
payload_dict["pattern"] = pattern
payload = TLVEncoder.encode(payload_dict)
header = BinaryProtocol.encode_header(MSG_TYPE_STASH, request_id,
len(payload))
return header + payload
@staticmethod
def decode_message(data: bytes) -> tuple:
"""
Decode a complete message (header + payload).
Args:
data: Complete message bytes
Returns:
tuple: (msg_type_str, request_id, payload_dict)
msg_type_str is the string name (e.g., "request")
payload_dict is the decoded TLV payload as a dict
Raises:
DirtyProtocolError: If message is malformed
"""
msg_type, request_id, length = BinaryProtocol.decode_header(data)
if len(data) < HEADER_SIZE + length:
raise DirtyProtocolError(
f"Incomplete message: expected {HEADER_SIZE + length} bytes, "
f"got {len(data)}",
raw_data=data[:50]
)
if length == 0:
# End message has empty payload
payload_dict = {}
else:
payload_data = data[HEADER_SIZE:HEADER_SIZE + length]
try:
payload_dict = TLVEncoder.decode_full(payload_data)
except DirtyProtocolError:
raise
except Exception as e:
raise DirtyProtocolError(
f"Failed to decode TLV payload: {e}",
raw_data=payload_data[:50]
)
# Convert to dict format similar to old JSON protocol
msg_type_str = MSG_TYPE_TO_STR[msg_type]
return msg_type_str, request_id, payload_dict
# -------------------------------------------------------------------------
# Async API (primary - for DirtyArbiter and DirtyWorker)
# -------------------------------------------------------------------------
@staticmethod
async def read_message_async(reader: asyncio.StreamReader) -> dict:
"""
Read a complete binary message from async stream.
Args:
reader: asyncio StreamReader
Returns:
dict: Message dict with 'type', 'id', and payload fields
Raises:
DirtyProtocolError: If read fails or message is malformed
asyncio.IncompleteReadError: If connection closed mid-read
"""
# Read header
try:
header = await reader.readexactly(HEADER_SIZE)
except asyncio.IncompleteReadError as e:
if len(e.partial) == 0:
# Clean close - no data was read
raise
raise DirtyProtocolError(
f"Incomplete header: got {len(e.partial)} bytes, "
f"expected {HEADER_SIZE}",
raw_data=e.partial
)
msg_type, request_id, length = BinaryProtocol.decode_header(header)
# Read payload
if length > 0:
try:
payload_data = await reader.readexactly(length)
except asyncio.IncompleteReadError as e:
raise DirtyProtocolError(
f"Incomplete payload: got {len(e.partial)} bytes, "
f"expected {length}",
raw_data=e.partial
)
try:
payload_dict = TLVEncoder.decode_full(payload_data)
except DirtyProtocolError:
raise
except Exception as e:
raise DirtyProtocolError(
f"Failed to decode TLV payload: {e}",
raw_data=payload_data[:50]
)
else:
payload_dict = {}
# Build response dict
msg_type_str = MSG_TYPE_TO_STR[msg_type]
result = {"type": msg_type_str, "id": request_id}
result.update(payload_dict)
return result
@staticmethod
async def write_message_async(writer: asyncio.StreamWriter,
message: dict) -> None:
"""
Write a message to async stream.
Accepts dict format for backwards compatibility.
Args:
writer: asyncio StreamWriter
message: Message dict with 'type', 'id', and payload fields
Raises:
DirtyProtocolError: If encoding fails
ConnectionError: If write fails
"""
data = BinaryProtocol._encode_from_dict(message)
writer.write(data)
await writer.drain()
# -------------------------------------------------------------------------
# Sync API (for HTTP workers that may not be async)
# -------------------------------------------------------------------------
@staticmethod
def _recv_exactly(sock: socket.socket, n: int) -> bytes:
"""
Receive exactly n bytes from a socket.
Args:
sock: Socket to read from
n: Number of bytes to read
Returns:
bytes: Received data
Raises:
DirtyProtocolError: If read fails or connection closed
"""
data = b""
while len(data) < n:
chunk = sock.recv(n - len(data))
if not chunk:
if len(data) == 0:
raise DirtyProtocolError("Connection closed")
raise DirtyProtocolError(
f"Connection closed after {len(data)} bytes, expected {n}",
raw_data=data
)
data += chunk
return data
@staticmethod
def read_message(sock: socket.socket) -> dict:
"""
Read a complete message from socket (sync).
Args:
sock: Socket to read from
Returns:
dict: Message dict with 'type', 'id', and payload fields
Raises:
DirtyProtocolError: If read fails or message is malformed
"""
# Read header
header = BinaryProtocol._recv_exactly(sock, HEADER_SIZE)
msg_type, request_id, length = BinaryProtocol.decode_header(header)
# Read payload
if length > 0:
payload_data = BinaryProtocol._recv_exactly(sock, length)
try:
payload_dict = TLVEncoder.decode_full(payload_data)
except DirtyProtocolError:
raise
except Exception as e:
raise DirtyProtocolError(
f"Failed to decode TLV payload: {e}",
raw_data=payload_data[:50]
)
else:
payload_dict = {}
# Build response dict
msg_type_str = MSG_TYPE_TO_STR[msg_type]
result = {"type": msg_type_str, "id": request_id}
result.update(payload_dict)
return result
@staticmethod
def write_message(sock: socket.socket, message: dict) -> None:
"""
Write a message to socket (sync).
Args:
sock: Socket to write to
message: Message dict with 'type', 'id', and payload fields
Raises:
DirtyProtocolError: If encoding fails
OSError: If write fails
"""
data = BinaryProtocol._encode_from_dict(message)
sock.sendall(data)
@staticmethod
def _encode_from_dict(message: dict) -> bytes: # pylint: disable=too-many-return-statements
"""
Encode a message dict to binary format.
Supports the old dict-based API for backwards compatibility.
Args:
message: Message dict with 'type', 'id', and payload fields
Returns:
bytes: Complete encoded message
"""
msg_type_str = message.get("type")
request_id = message.get("id", 0)
# Handle string or int request IDs
if isinstance(request_id, str):
# For backwards compat with UUID strings, hash to int
request_id = hash(request_id) & 0xFFFFFFFFFFFFFFFF
msg_type = MSG_TYPE_FROM_STR.get(msg_type_str)
if msg_type is None:
raise DirtyProtocolError(f"Unknown message type: {msg_type_str}")
if msg_type == MSG_TYPE_REQUEST:
return BinaryProtocol.encode_request(
request_id,
message.get("app_path", ""),
message.get("action", ""),
message.get("args"),
message.get("kwargs")
)
elif msg_type == MSG_TYPE_RESPONSE:
return BinaryProtocol.encode_response(
request_id,
message.get("result")
)
elif msg_type == MSG_TYPE_ERROR:
return BinaryProtocol.encode_error(
request_id,
message.get("error", {})
)
elif msg_type == MSG_TYPE_CHUNK:
return BinaryProtocol.encode_chunk(
request_id,
message.get("data")
)
elif msg_type == MSG_TYPE_END:
return BinaryProtocol.encode_end(request_id)
elif msg_type == MSG_TYPE_STASH:
return BinaryProtocol.encode_stash(
request_id,
message.get("op"),
message.get("table", ""),
message.get("key"),
message.get("value"),
message.get("pattern")
)
elif msg_type == MSG_TYPE_STATUS:
return BinaryProtocol.encode_status(request_id)
elif msg_type == MSG_TYPE_MANAGE:
return BinaryProtocol.encode_manage(
request_id,
message.get("op"),
message.get("count", 1)
)
else:
raise DirtyProtocolError(f"Unhandled message type: {msg_type}")
# =============================================================================
# Backwards Compatibility Aliases
# =============================================================================
# Alias BinaryProtocol as DirtyProtocol for drop-in replacement
DirtyProtocol = BinaryProtocol
# Message builder helpers (backwards compatible with old API)
def make_request(request_id, app_path: str, action: str,
args: tuple = None, kwargs: dict = None) -> dict:
"""
Build a request message dict.
Args:
request_id: Unique request identifier (int or str)
app_path: Import path of the dirty app (e.g., 'myapp.ml:MLApp')
action: Action to call on the app
args: Positional arguments
kwargs: Keyword arguments
Returns:
dict: Request message dict
"""
return {
"type": DirtyProtocol.MSG_TYPE_REQUEST,
"id": request_id,
"app_path": app_path,
"action": action,
"args": list(args) if args else [],
"kwargs": kwargs or {},
}
def make_response(request_id, result) -> dict:
"""
Build a success response message dict.
Args:
request_id: Request identifier this responds to
result: Result value
Returns:
dict: Response message dict
"""
return {
"type": DirtyProtocol.MSG_TYPE_RESPONSE,
"id": request_id,
"result": result,
}
def make_error_response(request_id, error) -> dict:
"""
Build an error response message dict.
Args:
request_id: Request identifier this responds to
error: DirtyError instance or dict with error info
Returns:
dict: Error response message dict
"""
from .errors import DirtyError
if isinstance(error, DirtyError):
error_dict = error.to_dict()
elif isinstance(error, dict):
error_dict = error
else:
error_dict = {
"error_type": type(error).__name__,
"message": str(error),
"details": {},
}
return {
"type": DirtyProtocol.MSG_TYPE_ERROR,
"id": request_id,
"error": error_dict,
}
def make_chunk_message(request_id, data) -> dict:
"""
Build a chunk message dict for streaming responses.
Args:
request_id: Request identifier this chunk belongs to
data: Chunk data
Returns:
dict: Chunk message dict
"""
return {
"type": DirtyProtocol.MSG_TYPE_CHUNK,
"id": request_id,
"data": data,
}
def make_end_message(request_id) -> dict:
"""
Build an end-of-stream message dict.
Args:
request_id: Request identifier this ends
Returns:
dict: End message dict
"""
return {
"type": DirtyProtocol.MSG_TYPE_END,
"id": request_id,
}
def make_stash_message(request_id, op: int, table: str,
key=None, value=None, pattern=None) -> dict:
"""
Build a stash operation message dict.
Args:
request_id: Unique request identifier (int or str)
op: Stash operation code (STASH_OP_*)
table: Table name
key: Optional key for put/get/delete operations
value: Optional value for put operation
pattern: Optional pattern for keys operation
Returns:
dict: Stash message dict
"""
msg = {
"type": DirtyProtocol.MSG_TYPE_STASH,
"id": request_id,
"op": op,
"table": table,
}
if key is not None:
msg["key"] = key
if value is not None:
msg["value"] = value
if pattern is not None:
msg["pattern"] = pattern
return msg
def make_manage_message(request_id, op: int, count: int = 1) -> dict:
"""
Build a worker management message dict.
Args:
request_id: Unique request identifier (int or str)
op: Management operation (MANAGE_OP_ADD or MANAGE_OP_REMOVE)
count: Number of workers to add/remove
Returns:
dict: Manage message dict
"""
return {
"type": DirtyProtocol.MSG_TYPE_MANAGE,
"id": request_id,
"op": op,
"count": count,
}
+503
View File
@@ -0,0 +1,503 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Stash - Global Shared State for Dirty Workers
Provides simple key-value tables stored in the arbiter process.
All workers can read and write to the same tables.
Usage::
from gunicorn.dirty import stash
# Basic operations - table is auto-created on first access
stash.put("sessions", "user:1", {"name": "Alice", "role": "admin"})
user = stash.get("sessions", "user:1")
stash.delete("sessions", "user:1")
# Dict-like interface
sessions = stash.table("sessions")
sessions["user:1"] = {"name": "Alice"}
user = sessions["user:1"]
del sessions["user:1"]
# Query operations
keys = stash.keys("sessions")
keys = stash.keys("sessions", pattern="user:*")
# Table management
stash.ensure("cache") # Explicit creation (idempotent)
stash.clear("sessions") # Delete all entries
stash.delete_table("sessions") # Delete the table itself
tables = stash.tables() # List all tables
Declarative usage in DirtyApp::
class MyApp(DirtyApp):
stashes = ["sessions", "cache"] # Auto-created on arbiter start
def __call__(self, action, *args, **kwargs):
# Tables are ready to use
stash.put("sessions", "key", "value")
Note: Tables are stored in the arbiter process and are ephemeral.
If the arbiter restarts, all data is lost.
"""
import threading
import uuid
from .errors import DirtyError
from .protocol import (
DirtyProtocol,
STASH_OP_PUT,
STASH_OP_GET,
STASH_OP_DELETE,
STASH_OP_KEYS,
STASH_OP_CLEAR,
STASH_OP_INFO,
STASH_OP_ENSURE,
STASH_OP_DELETE_TABLE,
STASH_OP_TABLES,
STASH_OP_EXISTS,
make_stash_message,
)
class StashError(DirtyError):
"""Base exception for stash operations."""
class StashTableNotFoundError(StashError):
"""Raised when a table does not exist."""
def __init__(self, table_name):
self.table_name = table_name
super().__init__(f"Stash table not found: {table_name}")
class StashKeyNotFoundError(StashError):
"""Raised when a key does not exist in a table."""
def __init__(self, table_name, key):
self.table_name = table_name
self.key = key
super().__init__(f"Key not found in {table_name}: {key}")
class StashClient:
"""
Client for stash operations.
Communicates with the arbiter which stores all tables in memory.
"""
def __init__(self, socket_path, timeout=30.0):
"""
Initialize the stash client.
Args:
socket_path: Path to the dirty arbiter's Unix socket
timeout: Default timeout for operations in seconds
"""
self.socket_path = socket_path
self.timeout = timeout
self._sock = None
self._lock = threading.Lock()
def _get_request_id(self):
"""Generate a unique request ID."""
return str(uuid.uuid4())
def _connect(self):
"""Establish connection to arbiter."""
import socket
if self._sock is not None:
return
try:
self._sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
self._sock.settimeout(self.timeout)
self._sock.connect(self.socket_path)
except (socket.error, OSError) as e:
self._sock = None
raise StashError(f"Failed to connect to arbiter: {e}") from e
def _close(self):
"""Close the connection."""
if self._sock is not None:
try:
self._sock.close()
except Exception:
pass
self._sock = None
def _execute(self, op, table, key=None, value=None, pattern=None):
"""
Execute a stash operation.
Args:
op: Operation code (STASH_OP_*)
table: Table name
key: Optional key
value: Optional value
pattern: Optional pattern for keys operation
Returns:
Result from the operation
"""
with self._lock:
if self._sock is None:
self._connect()
request_id = self._get_request_id()
message = make_stash_message(
request_id, op, table,
key=key, value=value, pattern=pattern
)
try:
DirtyProtocol.write_message(self._sock, message)
response = DirtyProtocol.read_message(self._sock)
msg_type = response.get("type")
if msg_type == DirtyProtocol.MSG_TYPE_RESPONSE:
return response.get("result")
elif msg_type == DirtyProtocol.MSG_TYPE_ERROR:
error_info = response.get("error", {})
error_type = error_info.get("error_type", "StashError")
error_msg = error_info.get("message", "Unknown error")
if error_type == "StashTableNotFoundError":
raise StashTableNotFoundError(table)
if error_type == "StashKeyNotFoundError":
raise StashKeyNotFoundError(table, key)
raise StashError(error_msg)
else:
raise StashError(f"Unexpected response type: {msg_type}")
except Exception as e:
self._close()
if isinstance(e, StashError):
raise
raise StashError(f"Stash operation failed: {e}") from e
# -------------------------------------------------------------------------
# Public API
# -------------------------------------------------------------------------
def put(self, table, key, value):
"""
Store a value in a table.
The table is automatically created if it doesn't exist.
Args:
table: Table name
key: Key to store under
value: Value to store (must be serializable)
"""
self._execute(STASH_OP_PUT, table, key=key, value=value)
def get(self, table, key, default=None):
"""
Retrieve a value from a table.
Args:
table: Table name
key: Key to retrieve
default: Default value if key not found
Returns:
The stored value, or default if not found
"""
try:
return self._execute(STASH_OP_GET, table, key=key)
except StashKeyNotFoundError:
return default
def delete(self, table, key):
"""
Delete a key from a table.
Args:
table: Table name
key: Key to delete
Returns:
True if key was deleted, False if it didn't exist
"""
return self._execute(STASH_OP_DELETE, table, key=key)
def keys(self, table, pattern=None):
"""
Get all keys in a table, optionally filtered by pattern.
Args:
table: Table name
pattern: Optional glob pattern (e.g., "user:*")
Returns:
List of keys
"""
return self._execute(STASH_OP_KEYS, table, pattern=pattern)
def clear(self, table):
"""
Delete all entries in a table.
Args:
table: Table name
"""
self._execute(STASH_OP_CLEAR, table)
def info(self, table):
"""
Get information about a table.
Args:
table: Table name
Returns:
Dict with table info (size, etc.)
"""
return self._execute(STASH_OP_INFO, table)
def ensure(self, table):
"""
Ensure a table exists (create if not exists).
This is idempotent - calling it multiple times is safe.
Args:
table: Table name
"""
self._execute(STASH_OP_ENSURE, table)
def exists(self, table, key=None):
"""
Check if a table or key exists.
Args:
table: Table name
key: Optional key to check within the table
Returns:
True if exists, False otherwise
"""
return self._execute(STASH_OP_EXISTS, table, key=key)
def delete_table(self, table):
"""
Delete an entire table.
Args:
table: Table name
"""
self._execute(STASH_OP_DELETE_TABLE, table)
def tables(self):
"""
List all tables.
Returns:
List of table names
"""
return self._execute(STASH_OP_TABLES, "")
def table(self, name):
"""
Get a dict-like interface to a table.
Args:
name: Table name
Returns:
StashTable instance
"""
return StashTable(self, name)
def close(self):
"""Close the client connection."""
with self._lock:
self._close()
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb):
self.close()
class StashTable:
"""
Dict-like interface to a stash table.
Example::
sessions = stash.table("sessions")
sessions["user:1"] = {"name": "Alice"}
user = sessions["user:1"]
del sessions["user:1"]
# Iteration
for key in sessions:
print(key, sessions[key])
"""
def __init__(self, client, name):
self._client = client
self._name = name
@property
def name(self):
"""Table name."""
return self._name
def __getitem__(self, key):
result = self._client.get(self._name, key)
if result is None:
# Check if key actually exists with None value
if not self._client.exists(self._name, key):
raise KeyError(key)
return result
def __setitem__(self, key, value):
self._client.put(self._name, key, value)
def __delitem__(self, key):
if not self._client.delete(self._name, key):
raise KeyError(key)
def __contains__(self, key):
return self._client.exists(self._name, key)
def __iter__(self):
return iter(self._client.keys(self._name))
def __len__(self):
info = self._client.info(self._name)
return info.get("size", 0)
def get(self, key, default=None):
"""Get value with default."""
return self._client.get(self._name, key, default)
def keys(self, pattern=None):
"""Get all keys, optionally filtered by pattern."""
return self._client.keys(self._name, pattern=pattern)
def clear(self):
"""Delete all entries."""
self._client.clear(self._name)
def items(self):
"""Iterate over (key, value) pairs."""
for key in self._client.keys(self._name):
yield key, self._client.get(self._name, key)
def values(self):
"""Iterate over values."""
for key in self._client.keys(self._name):
yield self._client.get(self._name, key)
# =============================================================================
# Global stash instance (module-level API)
# =============================================================================
# Thread-local storage for stash clients
_thread_local = threading.local()
# Global socket path
_stash_socket_path = None
def set_stash_socket_path(path):
"""Set the global stash socket path (called during initialization)."""
global _stash_socket_path # pylint: disable=global-statement
_stash_socket_path = path
def get_stash_socket_path():
"""Get the stash socket path."""
import os
if _stash_socket_path is None:
# Check environment variable
path = os.environ.get('GUNICORN_DIRTY_SOCKET')
if path:
return path
raise StashError(
"Stash socket path not configured. "
"Make sure dirty_workers > 0 and dirty_apps are configured."
)
return _stash_socket_path
def _get_client():
"""Get or create a thread-local stash client."""
client = getattr(_thread_local, 'stash_client', None)
if client is None:
socket_path = get_stash_socket_path()
client = StashClient(socket_path)
_thread_local.stash_client = client
return client
# Module-level functions that use the thread-local client
def put(table, key, value):
"""Store a value in a table."""
_get_client().put(table, key, value)
def get(table, key, default=None):
"""Retrieve a value from a table."""
return _get_client().get(table, key, default)
def delete(table, key):
"""Delete a key from a table."""
return _get_client().delete(table, key)
def keys(table, pattern=None):
"""Get all keys in a table."""
return _get_client().keys(table, pattern)
def clear(table):
"""Delete all entries in a table."""
_get_client().clear(table)
def info(table):
"""Get information about a table."""
return _get_client().info(table)
def ensure(table):
"""Ensure a table exists."""
_get_client().ensure(table)
def exists(table, key=None):
"""Check if a table or key exists."""
return _get_client().exists(table, key)
def delete_table(table):
"""Delete an entire table."""
_get_client().delete_table(table)
def tables():
"""List all tables."""
return _get_client().tables()
def table(name):
"""Get a dict-like interface to a table."""
return _get_client().table(name)
+303
View File
@@ -0,0 +1,303 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
TLV (Type-Length-Value) Binary Encoder/Decoder
Provides efficient binary serialization for dirty worker protocol messages.
Inspired by OpenBSD msgctl/msgsnd message format.
Type Codes:
0x00: None (no value bytes)
0x01: bool (1 byte: 0x00 or 0x01)
0x05: int64 (8 bytes big-endian signed)
0x06: float64 (8 bytes IEEE 754)
0x10: bytes (4-byte length + raw bytes)
0x11: string (4-byte length + UTF-8 encoded)
0x20: list (4-byte count + encoded elements)
0x21: dict (4-byte count + encoded key-value pairs)
"""
import struct
from .errors import DirtyProtocolError
# Type codes
TYPE_NONE = 0x00
TYPE_BOOL = 0x01
TYPE_INT64 = 0x05
TYPE_FLOAT64 = 0x06
TYPE_BYTES = 0x10
TYPE_STRING = 0x11
TYPE_LIST = 0x20
TYPE_DICT = 0x21
# Maximum sizes for safety
MAX_STRING_SIZE = 64 * 1024 * 1024 # 64 MB
MAX_BYTES_SIZE = 64 * 1024 * 1024 # 64 MB
MAX_LIST_SIZE = 1024 * 1024 # 1 million items
MAX_DICT_SIZE = 1024 * 1024 # 1 million items
class TLVEncoder:
"""
TLV binary encoder/decoder.
Encodes Python values to binary TLV format and decodes back.
Supports: None, bool, int, float, bytes, str, list, dict.
"""
@staticmethod
def encode(value) -> bytes: # pylint: disable=too-many-return-statements
"""
Encode a Python value to TLV binary format.
Args:
value: Python value to encode (None, bool, int, float,
bytes, str, list, or dict)
Returns:
bytes: TLV-encoded binary data
Raises:
DirtyProtocolError: If value type is not supported
"""
if value is None:
return bytes([TYPE_NONE])
if isinstance(value, bool):
# bool must come before int since bool is a subclass of int
return bytes([TYPE_BOOL, 0x01 if value else 0x00])
if isinstance(value, int):
return bytes([TYPE_INT64]) + struct.pack(">q", value)
if isinstance(value, float):
return bytes([TYPE_FLOAT64]) + struct.pack(">d", value)
if isinstance(value, bytes):
if len(value) > MAX_BYTES_SIZE:
raise DirtyProtocolError(
f"Bytes too large: {len(value)} bytes "
f"(max: {MAX_BYTES_SIZE})"
)
return bytes([TYPE_BYTES]) + struct.pack(">I", len(value)) + value
if isinstance(value, str):
encoded = value.encode("utf-8")
if len(encoded) > MAX_STRING_SIZE:
raise DirtyProtocolError(
f"String too large: {len(encoded)} bytes "
f"(max: {MAX_STRING_SIZE})"
)
return bytes([TYPE_STRING]) + struct.pack(">I", len(encoded)) + encoded
if isinstance(value, (list, tuple)):
if len(value) > MAX_LIST_SIZE:
raise DirtyProtocolError(
f"List too large: {len(value)} items "
f"(max: {MAX_LIST_SIZE})"
)
parts = [bytes([TYPE_LIST]), struct.pack(">I", len(value))]
for item in value:
parts.append(TLVEncoder.encode(item))
return b"".join(parts)
if isinstance(value, dict):
if len(value) > MAX_DICT_SIZE:
raise DirtyProtocolError(
f"Dict too large: {len(value)} items "
f"(max: {MAX_DICT_SIZE})"
)
parts = [bytes([TYPE_DICT]), struct.pack(">I", len(value))]
for k, v in value.items():
# Convert keys to strings (like JSON)
if not isinstance(k, str):
k = str(k)
parts.append(TLVEncoder.encode(k))
parts.append(TLVEncoder.encode(v))
return b"".join(parts)
raise DirtyProtocolError(
f"Unsupported type for TLV encoding: {type(value).__name__}"
)
@staticmethod
def decode(data: bytes, offset: int = 0) -> tuple: # pylint: disable=too-many-return-statements
"""
Decode a TLV-encoded value from binary data.
Args:
data: Binary data to decode
offset: Starting offset in the data
Returns:
tuple: (decoded_value, new_offset)
Raises:
DirtyProtocolError: If data is malformed or truncated
"""
if offset >= len(data):
raise DirtyProtocolError(
"Truncated TLV data: no type byte",
raw_data=data[offset:offset + 20]
)
type_code = data[offset]
offset += 1
if type_code == TYPE_NONE:
return None, offset
if type_code == TYPE_BOOL:
if offset >= len(data):
raise DirtyProtocolError(
"Truncated TLV data: missing bool value",
raw_data=data[offset - 1:offset + 20]
)
value = data[offset] != 0x00
return value, offset + 1
if type_code == TYPE_INT64:
if offset + 8 > len(data):
raise DirtyProtocolError(
"Truncated TLV data: incomplete int64",
raw_data=data[offset - 1:offset + 20]
)
value = struct.unpack(">q", data[offset:offset + 8])[0]
return value, offset + 8
if type_code == TYPE_FLOAT64:
if offset + 8 > len(data):
raise DirtyProtocolError(
"Truncated TLV data: incomplete float64",
raw_data=data[offset - 1:offset + 20]
)
value = struct.unpack(">d", data[offset:offset + 8])[0]
return value, offset + 8
if type_code == TYPE_BYTES:
if offset + 4 > len(data):
raise DirtyProtocolError(
"Truncated TLV data: incomplete bytes length",
raw_data=data[offset - 1:offset + 20]
)
length = struct.unpack(">I", data[offset:offset + 4])[0]
offset += 4
if length > MAX_BYTES_SIZE:
raise DirtyProtocolError(
f"Bytes too large: {length} bytes (max: {MAX_BYTES_SIZE})"
)
if offset + length > len(data):
raise DirtyProtocolError(
f"Truncated TLV data: expected {length} bytes, "
f"got {len(data) - offset}",
raw_data=data[offset - 5:offset + 20]
)
value = data[offset:offset + length]
return value, offset + length
if type_code == TYPE_STRING:
if offset + 4 > len(data):
raise DirtyProtocolError(
"Truncated TLV data: incomplete string length",
raw_data=data[offset - 1:offset + 20]
)
length = struct.unpack(">I", data[offset:offset + 4])[0]
offset += 4
if length > MAX_STRING_SIZE:
raise DirtyProtocolError(
f"String too large: {length} bytes (max: {MAX_STRING_SIZE})"
)
if offset + length > len(data):
raise DirtyProtocolError(
f"Truncated TLV data: expected {length} bytes for string, "
f"got {len(data) - offset}",
raw_data=data[offset - 5:offset + 20]
)
try:
value = data[offset:offset + length].decode("utf-8")
except UnicodeDecodeError as e:
raise DirtyProtocolError(
f"Invalid UTF-8 in string: {e}",
raw_data=data[offset:offset + min(length, 20)]
)
return value, offset + length
if type_code == TYPE_LIST:
if offset + 4 > len(data):
raise DirtyProtocolError(
"Truncated TLV data: incomplete list count",
raw_data=data[offset - 1:offset + 20]
)
count = struct.unpack(">I", data[offset:offset + 4])[0]
offset += 4
if count > MAX_LIST_SIZE:
raise DirtyProtocolError(
f"List too large: {count} items (max: {MAX_LIST_SIZE})"
)
items = []
for _ in range(count):
item, offset = TLVEncoder.decode(data, offset)
items.append(item)
return items, offset
if type_code == TYPE_DICT:
if offset + 4 > len(data):
raise DirtyProtocolError(
"Truncated TLV data: incomplete dict count",
raw_data=data[offset - 1:offset + 20]
)
count = struct.unpack(">I", data[offset:offset + 4])[0]
offset += 4
if count > MAX_DICT_SIZE:
raise DirtyProtocolError(
f"Dict too large: {count} items (max: {MAX_DICT_SIZE})"
)
result = {}
for _ in range(count):
key, offset = TLVEncoder.decode(data, offset)
if not isinstance(key, str):
raise DirtyProtocolError(
f"Dict key must be string, got {type(key).__name__}"
)
value, offset = TLVEncoder.decode(data, offset)
result[key] = value
return result, offset
raise DirtyProtocolError(
f"Unknown TLV type code: 0x{type_code:02x}",
raw_data=data[offset - 1:offset + 20]
)
@staticmethod
def decode_full(data: bytes):
"""
Decode a complete TLV-encoded value, ensuring all data is consumed.
Args:
data: Binary data to decode
Returns:
Decoded Python value
Raises:
DirtyProtocolError: If data is malformed or has trailing bytes
"""
value, offset = TLVEncoder.decode(data, 0)
if offset != len(data):
raise DirtyProtocolError(
f"Trailing data after TLV: {len(data) - offset} bytes",
raw_data=data[offset:offset + 20]
)
return value
+530
View File
@@ -0,0 +1,530 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Dirty Worker Process
Asyncio-based worker that loads dirty apps and handles requests
from the DirtyArbiter.
Threading Model
---------------
Each dirty worker runs an asyncio event loop in the main thread for:
- Handling connections from the arbiter
- Managing heartbeat updates
- Coordinating task execution
Actual app execution runs in a ThreadPoolExecutor (separate threads):
- The number of threads is controlled by ``dirty_threads`` config (default: 1)
- Each thread can execute one app action at a time
- The asyncio event loop is NOT blocked by task execution
State and Global Objects
------------------------
Apps can maintain persistent state because:
1. Apps are loaded ONCE when the worker starts (in ``load_apps()``)
2. The same app instances are reused for ALL requests
3. App state (instance variables, loaded models, etc.) persists
Example::
class MLApp(DirtyApp):
def init(self):
self.model = load_heavy_model() # Loaded once, reused
self.cache = {} # Persistent cache
def predict(self, data):
return self.model.predict(data) # Uses loaded model
Thread Safety:
- With ``dirty_threads=1`` (default): No concurrent access, thread-safe by design
- With ``dirty_threads > 1``: Multiple threads share the same app instances,
apps MUST be thread-safe (use locks, thread-local storage, etc.)
Heartbeat and Liveness
----------------------
The worker sends heartbeat updates to prove it's alive:
1. A dedicated asyncio task (``_heartbeat_loop``) runs independently
2. It updates the heartbeat file every ``dirty_timeout / 2`` seconds
3. Since tasks run in executor threads, they do NOT block heartbeats
4. The arbiter kills workers that miss heartbeat updates
Timeout Control
---------------
Execution timeout is enforced at two levels:
1. **Worker level**: Each task execution has a timeout (``dirty_timeout``).
If exceeded, the worker returns a timeout error but the thread may
continue running (Python threads cannot be cancelled).
2. **Arbiter level**: The arbiter also enforces timeout when waiting
for worker response. Workers that don't respond are killed via SIGABRT.
Note: Since Python threads cannot be forcibly cancelled, a truly stuck
operation will continue until the worker is killed by the arbiter.
"""
import asyncio
import inspect
import os
import signal
import traceback
import uuid
from gunicorn import util
from gunicorn.workers.workertmp import WorkerTmp
from .app import load_dirty_apps
from .errors import (
DirtyAppError,
DirtyAppNotFoundError,
DirtyTimeoutError,
DirtyWorkerError,
)
from .protocol import (
DirtyProtocol,
make_response,
make_error_response,
make_chunk_message,
make_end_message,
)
class DirtyWorker:
"""
Dirty worker process that loads dirty apps and handles requests.
Each worker runs its own asyncio event loop and listens on a
worker-specific Unix socket for requests from the DirtyArbiter.
"""
SIGNALS = [getattr(signal, "SIG%s" % x) for x in
"ABRT HUP QUIT INT TERM USR1".split()]
def __init__(self, age, ppid, app_paths, cfg, log, socket_path):
"""
Initialize a dirty worker.
Args:
age: Worker age (for identifying workers)
ppid: Parent process ID
app_paths: List of dirty app import paths
cfg: Gunicorn config
log: Logger
socket_path: Path to this worker's Unix socket
"""
self.age = age
self.pid = "[booting]"
self.ppid = ppid
self.app_paths = app_paths
self.cfg = cfg
self.log = log
self.socket_path = socket_path
self.booted = False
self.aborted = False
self.alive = True
self.tmp = WorkerTmp(cfg)
self.apps = {}
self._server = None
self._loop = None
self._executor = None
def __str__(self):
return f"<DirtyWorker {self.pid}>"
def notify(self):
"""Update heartbeat timestamp."""
self.tmp.notify()
def init_process(self):
"""
Initialize the worker process after fork.
This is called in the child process after fork. It sets up
the environment, loads apps, and starts the main run loop.
"""
# Set environment variables
if self.cfg.env:
for k, v in self.cfg.env.items():
os.environ[k] = v
util.set_owner_process(self.cfg.uid, self.cfg.gid,
initgroups=self.cfg.initgroups)
# Reseed random number generator
util.seed()
# Prevent fd inheritance
util.close_on_exec(self.tmp.fileno())
self.log.close_on_exec()
# Set up signals
self.init_signals()
# Load dirty apps
self.load_apps()
# Call hook
self.pid = os.getpid()
self.cfg.dirty_worker_init(self)
# Enter main run loop
self.booted = True
self.run()
def init_signals(self):
"""Set up signal handlers."""
# Reset signal handlers from parent
for sig in self.SIGNALS:
signal.signal(sig, signal.SIG_DFL)
# Handle graceful shutdown
signal.signal(signal.SIGTERM, self._signal_handler)
signal.signal(signal.SIGQUIT, self._signal_handler)
signal.signal(signal.SIGINT, self._signal_handler)
# Handle abort (timeout)
signal.signal(signal.SIGABRT, self._signal_handler)
# Handle USR1 (reopen logs)
signal.signal(signal.SIGUSR1, self._signal_handler)
def _signal_handler(self, sig, frame):
"""Handle signals by setting alive = False."""
if sig == signal.SIGUSR1:
self.log.reopen_files()
return
self.alive = False
if self._loop:
self._loop.call_soon_threadsafe(self._shutdown)
def _shutdown(self):
"""Initiate async shutdown."""
if self._server:
self._server.close()
def load_apps(self):
"""Load all configured dirty apps."""
try:
self.apps = load_dirty_apps(self.app_paths)
for path, app in self.apps.items():
self.log.debug("Loaded dirty app: %s", path)
try:
app.init()
self.log.info("Initialized dirty app: %s", path)
except Exception as e:
self.log.error("Failed to initialize dirty app %s: %s",
path, e)
raise
except Exception as e:
self.log.error("Failed to load dirty apps: %s", e)
raise
def run(self):
"""Run the main asyncio event loop."""
# Lazy import for gevent compatibility (see #3482)
from concurrent.futures import ThreadPoolExecutor
# Create thread pool for executing app actions
num_threads = self.cfg.dirty_threads
self._executor = ThreadPoolExecutor(
max_workers=num_threads,
thread_name_prefix=f"dirty-worker-{self.pid}-"
)
self.log.debug("Created thread pool with %d threads", num_threads)
try:
self._loop = asyncio.new_event_loop()
asyncio.set_event_loop(self._loop)
self._loop.run_until_complete(self._run_async())
except Exception as e:
self.log.error("Worker error: %s", e)
finally:
self._cleanup()
async def _run_async(self):
"""Main async loop - start server and handle connections."""
# Remove socket if it exists
if os.path.exists(self.socket_path):
os.unlink(self.socket_path)
# Start Unix socket server
self._server = await asyncio.start_unix_server(
self.handle_connection,
path=self.socket_path
)
# Make socket accessible
os.chmod(self.socket_path, 0o600)
self.log.info("Dirty worker %s listening on %s",
self.pid, self.socket_path)
# Start heartbeat task
heartbeat_task = asyncio.create_task(self._heartbeat_loop())
try:
async with self._server:
await self._server.serve_forever()
except asyncio.CancelledError:
pass
finally:
heartbeat_task.cancel()
try:
await heartbeat_task
except asyncio.CancelledError:
pass
async def _heartbeat_loop(self):
"""Periodically update heartbeat."""
while self.alive:
self.notify()
await asyncio.sleep(self.cfg.dirty_timeout / 2.0)
async def handle_connection(self, reader, writer):
"""
Handle a connection from the arbiter.
Each connection can send multiple requests.
"""
self.log.debug("New connection from arbiter")
try:
while self.alive:
try:
message = await DirtyProtocol.read_message_async(reader)
except asyncio.IncompleteReadError:
# Connection closed
break
# Handle the request - pass writer for streaming support
await self.handle_request(message, writer)
except Exception as e:
self.log.error("Connection error: %s", e)
finally:
writer.close()
try:
await writer.wait_closed()
except Exception:
pass
async def handle_request(self, message, writer):
"""
Handle a single request message.
Supports both regular (non-streaming) and streaming responses.
For streaming, detects if the result is a generator and sends
chunk messages followed by an end message.
Args:
message: Request dict from protocol
writer: StreamWriter for sending responses
"""
request_id = message.get("id", str(uuid.uuid4()))
msg_type = message.get("type")
if msg_type != DirtyProtocol.MSG_TYPE_REQUEST:
response = make_error_response(
request_id,
DirtyWorkerError(f"Unknown message type: {msg_type}")
)
await DirtyProtocol.write_message_async(writer, response)
return
app_path = message.get("app_path")
action = message.get("action")
args = message.get("args", [])
kwargs = message.get("kwargs", {})
# Update heartbeat before executing
self.notify()
try:
result = await self.execute(app_path, action, args, kwargs)
# Check if result is a generator (streaming)
if inspect.isgenerator(result):
await self._stream_sync_generator(request_id, result, writer)
elif inspect.isasyncgen(result):
await self._stream_async_generator(request_id, result, writer)
else:
# Regular non-streaming response
response = make_response(request_id, result)
await DirtyProtocol.write_message_async(writer, response)
except Exception as e:
tb = traceback.format_exc()
self.log.error("Error executing %s.%s: %s\n%s",
app_path, action, e, tb)
response = make_error_response(
request_id,
DirtyAppError(str(e), app_path=app_path, action=action,
traceback=tb)
)
await DirtyProtocol.write_message_async(writer, response)
async def _stream_sync_generator(self, request_id, gen, writer):
"""
Stream chunks from a synchronous generator.
Args:
request_id: Request ID for the messages
gen: Sync generator to iterate
writer: StreamWriter for sending messages
"""
# Sentinel value to detect end of generator
# (StopIteration cannot be raised into a Future in Python 3.7+)
_EXHAUSTED = object()
def _get_next():
try:
return next(gen)
except StopIteration:
return _EXHAUSTED
try:
loop = asyncio.get_running_loop()
while True:
# Run next() in executor to avoid blocking event loop
chunk = await loop.run_in_executor(self._executor, _get_next)
if chunk is _EXHAUSTED:
break
# Send chunk message
await DirtyProtocol.write_message_async(
writer, make_chunk_message(request_id, chunk)
)
# Update heartbeat during long streams
self.notify()
# Send end message
await DirtyProtocol.write_message_async(
writer, make_end_message(request_id)
)
except Exception as e:
# Error during streaming - send error message
tb = traceback.format_exc()
self.log.error("Error during streaming: %s\n%s", e, tb)
response = make_error_response(
request_id,
DirtyAppError(str(e), traceback=tb)
)
await DirtyProtocol.write_message_async(writer, response)
finally:
gen.close()
async def _stream_async_generator(self, request_id, gen, writer):
"""
Stream chunks from an asynchronous generator.
Args:
request_id: Request ID for the messages
gen: Async generator to iterate
writer: StreamWriter for sending messages
"""
try:
async for chunk in gen:
# Send chunk message
await DirtyProtocol.write_message_async(
writer, make_chunk_message(request_id, chunk)
)
# Update heartbeat during long streams
self.notify()
# Send end message
await DirtyProtocol.write_message_async(
writer, make_end_message(request_id)
)
except Exception as e:
# Error during streaming - send error message
tb = traceback.format_exc()
self.log.error("Error during streaming: %s\n%s", e, tb)
response = make_error_response(
request_id,
DirtyAppError(str(e), traceback=tb)
)
await DirtyProtocol.write_message_async(writer, response)
finally:
await gen.aclose()
async def execute(self, app_path, action, args, kwargs):
"""
Execute an action on a dirty app.
The action runs in a thread pool executor to avoid blocking the
asyncio event loop. Execution timeout is enforced using
``dirty_timeout`` config.
Args:
app_path: Import path of the dirty app
action: Action name to execute
args: Positional arguments
kwargs: Keyword arguments
Returns:
Result from the app action
Raises:
DirtyAppNotFoundError: If app is not loaded
DirtyTimeoutError: If execution exceeds timeout
DirtyAppError: If execution fails
"""
if app_path not in self.apps:
raise DirtyAppNotFoundError(app_path)
app = self.apps[app_path]
timeout = self.cfg.dirty_timeout if self.cfg.dirty_timeout > 0 else None
# Run the app call in the thread pool to avoid blocking
# the event loop for CPU-bound operations
loop = asyncio.get_running_loop()
try:
result = await asyncio.wait_for(
loop.run_in_executor(
self._executor,
lambda: app(action, *args, **kwargs)
),
timeout=timeout
)
return result
except asyncio.TimeoutError:
# Note: The thread continues running - we just stop waiting
self.log.warning(
"Execution timeout for %s.%s after %ds",
app_path, action, timeout
)
raise DirtyTimeoutError(
f"Execution of {app_path}.{action} timed out",
timeout=timeout
)
def _cleanup(self):
"""Clean up resources on shutdown."""
# Shutdown thread pool executor
if self._executor:
self._executor.shutdown(wait=False, cancel_futures=True)
self._executor = None
# Close all apps
for path, app in self.apps.items():
try:
app.close()
self.log.debug("Closed dirty app: %s", path)
except Exception as e:
self.log.error("Error closing dirty app %s: %s", path, e)
# Close temp file
try:
self.tmp.close()
except Exception:
pass
# Remove socket file
try:
if os.path.exists(self.socket_path):
os.unlink(self.socket_path)
except Exception:
pass
self.log.info("Dirty worker %s exiting", self.pid)
+28
View File
@@ -0,0 +1,28 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
# We don't need to call super() in __init__ methods of our
# BaseException and Exception classes because we also define
# our own __str__ methods so there is no need to pass 'message'
# to the base class to get a meaningful output from 'str(exc)'.
# pylint: disable=super-init-not-called
# we inherit from BaseException here to make sure to not be caught
# at application level
class HaltServer(BaseException):
def __init__(self, reason, exit_status=1):
self.reason = reason
self.exit_status = exit_status
def __str__(self):
return "<HaltServer %r %d>" % (self.reason, self.exit_status)
class ConfigError(Exception):
""" Exception raised on config error """
class AppImportError(Exception):
""" Exception raised when loading an application """
+493
View File
@@ -0,0 +1,493 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import base64
import binascii
import json
import time
import logging
logging.Logger.manager.emittedNoHandlerWarning = 1 # noqa
from logging.config import dictConfig
from logging.config import fileConfig
import os
import socket
import sys
import threading
import traceback
from gunicorn import util
# syslog facility codes
SYSLOG_FACILITIES = {
"auth": 4,
"authpriv": 10,
"cron": 9,
"daemon": 3,
"ftp": 11,
"kern": 0,
"lpr": 6,
"mail": 2,
"news": 7,
"security": 4, # DEPRECATED
"syslog": 5,
"user": 1,
"uucp": 8,
"local0": 16,
"local1": 17,
"local2": 18,
"local3": 19,
"local4": 20,
"local5": 21,
"local6": 22,
"local7": 23
}
CONFIG_DEFAULTS = {
"version": 1,
"disable_existing_loggers": False,
"root": {"level": "INFO", "handlers": ["console"]},
"loggers": {
"gunicorn.error": {
"level": "INFO",
"handlers": ["error_console"],
"propagate": True,
"qualname": "gunicorn.error"
},
"gunicorn.access": {
"level": "INFO",
"handlers": ["console"],
"propagate": True,
"qualname": "gunicorn.access"
}
},
"handlers": {
"console": {
"class": "logging.StreamHandler",
"formatter": "generic",
"stream": "ext://sys.stdout"
},
"error_console": {
"class": "logging.StreamHandler",
"formatter": "generic",
"stream": "ext://sys.stderr"
},
},
"formatters": {
"generic": {
"format": "%(asctime)s [%(process)d] [%(levelname)s] %(message)s",
"datefmt": "[%Y-%m-%d %H:%M:%S %z]",
"class": "logging.Formatter"
}
}
}
def loggers():
""" get list of all loggers """
root = logging.root
existing = list(root.manager.loggerDict.keys())
return [logging.getLogger(name) for name in existing]
class SafeAtoms(dict):
def __init__(self, atoms):
dict.__init__(self)
for key, value in atoms.items():
if isinstance(value, str):
self[key] = value.replace('"', '\\"')
else:
self[key] = value
def __getitem__(self, k):
if k.startswith("{"):
kl = k.lower()
if kl in self:
return super().__getitem__(kl)
else:
return "-"
if k in self:
return super().__getitem__(k)
else:
return '-'
def parse_syslog_address(addr):
# unix domain socket type depends on backend
# SysLogHandler will try both when given None
if addr.startswith("unix://"):
sock_type = None
# set socket type only if explicitly requested
parts = addr.split("#", 1)
if len(parts) == 2:
addr = parts[0]
if parts[1] == "dgram":
sock_type = socket.SOCK_DGRAM
return (sock_type, addr.split("unix://")[1])
if addr.startswith("udp://"):
addr = addr.split("udp://")[1]
socktype = socket.SOCK_DGRAM
elif addr.startswith("tcp://"):
addr = addr.split("tcp://")[1]
socktype = socket.SOCK_STREAM
else:
raise RuntimeError("invalid syslog address")
if '[' in addr and ']' in addr:
host = addr.split(']')[0][1:].lower()
elif ':' in addr:
host = addr.split(':')[0].lower()
elif addr == "":
host = "localhost"
else:
host = addr.lower()
addr = addr.split(']')[-1]
if ":" in addr:
port = addr.split(':', 1)[1]
if not port.isdigit():
raise RuntimeError("%r is not a valid port number." % port)
port = int(port)
else:
port = 514
return (socktype, (host, port))
class Logger:
LOG_LEVELS = {
"critical": logging.CRITICAL,
"error": logging.ERROR,
"warning": logging.WARNING,
"info": logging.INFO,
"debug": logging.DEBUG
}
loglevel = logging.INFO
error_fmt = r"%(asctime)s [%(process)d] [%(levelname)s] %(message)s"
datefmt = r"[%Y-%m-%d %H:%M:%S %z]"
access_fmt = "%(message)s"
syslog_fmt = "[%(process)d] %(message)s"
atoms_wrapper_class = SafeAtoms
def __init__(self, cfg):
self.error_log = logging.getLogger("gunicorn.error")
self.error_log.propagate = False
self.access_log = logging.getLogger("gunicorn.access")
self.access_log.propagate = False
self.error_handlers = []
self.access_handlers = []
self.logfile = None
self.lock = threading.Lock()
self.cfg = cfg
self.setup(cfg)
def setup(self, cfg):
# the cfg object is replaced on each configuration reload, so
# follow the one we are given instead of the startup one
self.cfg = cfg
self.loglevel = self.LOG_LEVELS.get(cfg.loglevel.lower(), logging.INFO)
self.error_log.setLevel(self.loglevel)
self.access_log.setLevel(logging.INFO)
# drop our own handlers from a previous setup so that re-running
# it on configuration reload does not stack duplicates
for log in (self.error_log, self.access_log):
for handler in list(log.handlers):
if getattr(handler, "_gunicorn", False):
log.removeHandler(handler)
# set gunicorn.error handler
if self.cfg.capture_output and cfg.errorlog != "-":
for stream in sys.stdout, sys.stderr:
stream.flush()
self.logfile = open(cfg.errorlog, 'a+')
os.dup2(self.logfile.fileno(), sys.stdout.fileno())
os.dup2(self.logfile.fileno(), sys.stderr.fileno())
self._set_handler(self.error_log, cfg.errorlog,
logging.Formatter(self.error_fmt, self.datefmt))
# set gunicorn.access handler
if cfg.accesslog is not None:
self._set_handler(
self.access_log, cfg.accesslog,
fmt=logging.Formatter(self.access_fmt), stream=sys.stdout
)
# set syslog handler
if cfg.syslog:
self._set_syslog_handler(
self.error_log, cfg, self.syslog_fmt, "error"
)
if not cfg.disable_redirect_access_to_syslog:
self._set_syslog_handler(
self.access_log, cfg, self.syslog_fmt, "access"
)
if cfg.logconfig_dict:
config = CONFIG_DEFAULTS.copy()
config.update(cfg.logconfig_dict)
try:
dictConfig(config)
except (
AttributeError,
ImportError,
ValueError,
TypeError
) as exc:
raise RuntimeError(str(exc)) from exc
elif cfg.logconfig_json:
config = CONFIG_DEFAULTS.copy()
if os.path.exists(cfg.logconfig_json):
try:
config_json = json.load(open(cfg.logconfig_json))
config.update(config_json)
dictConfig(config)
except (
json.JSONDecodeError,
AttributeError,
ImportError,
ValueError,
TypeError
) as exc:
raise RuntimeError(str(exc)) from exc
elif cfg.logconfig:
if os.path.exists(cfg.logconfig):
defaults = CONFIG_DEFAULTS.copy()
defaults['__file__'] = cfg.logconfig
defaults['here'] = os.path.dirname(cfg.logconfig)
fileConfig(cfg.logconfig, defaults=defaults,
disable_existing_loggers=False)
else:
msg = "Error: log config '%s' not found"
raise RuntimeError(msg % cfg.logconfig)
def critical(self, msg, *args, **kwargs):
self.error_log.critical(msg, *args, **kwargs)
def error(self, msg, *args, **kwargs):
self.error_log.error(msg, *args, **kwargs)
def warning(self, msg, *args, **kwargs):
self.error_log.warning(msg, *args, **kwargs)
def info(self, msg, *args, **kwargs):
self.error_log.info(msg, *args, **kwargs)
def debug(self, msg, *args, **kwargs):
self.error_log.debug(msg, *args, **kwargs)
def exception(self, msg, *args, **kwargs):
self.error_log.exception(msg, *args, **kwargs)
def log(self, lvl, msg, *args, **kwargs):
if isinstance(lvl, str):
lvl = self.LOG_LEVELS.get(lvl.lower(), logging.INFO)
self.error_log.log(lvl, msg, *args, **kwargs)
def atoms(self, resp, req, environ, request_time):
""" Gets atoms for log formatting.
"""
status = resp.status
if isinstance(status, str):
status = status.split(None, 1)[0]
atoms = {
'h': environ.get('REMOTE_ADDR', '-'),
'l': '-',
'u': self._get_user(environ) or '-',
't': self.now(),
'r': "%s %s %s" % (environ['REQUEST_METHOD'],
environ['RAW_URI'],
environ["SERVER_PROTOCOL"]),
's': status,
'm': environ.get('REQUEST_METHOD'),
'U': environ.get('PATH_INFO'),
'q': environ.get('QUERY_STRING'),
'H': environ.get('SERVER_PROTOCOL'),
'b': getattr(resp, 'sent', None) is not None and str(resp.sent) or '-',
'B': getattr(resp, 'sent', None),
'f': environ.get('HTTP_REFERER', '-'),
'a': environ.get('HTTP_USER_AGENT', '-'),
'T': request_time.seconds,
'D': (request_time.seconds * 1000000) + request_time.microseconds,
'M': (request_time.seconds * 1000) + int(request_time.microseconds / 1000),
'L': "%d.%06d" % (request_time.seconds, request_time.microseconds),
'p': "<%s>" % os.getpid()
}
# add request headers
if hasattr(req, 'headers'):
req_headers = req.headers
else:
req_headers = req
if hasattr(req_headers, "items"):
req_headers = req_headers.items()
atoms.update({"{%s}i" % k.lower(): v for k, v in req_headers})
resp_headers = resp.headers
if hasattr(resp_headers, "items"):
resp_headers = resp_headers.items()
# add response headers
atoms.update({"{%s}o" % k.lower(): v for k, v in resp_headers})
# add environ variables
environ_variables = environ.items()
atoms.update({"{%s}e" % k.lower(): v for k, v in environ_variables})
return atoms
@property
def access_log_enabled(self):
"""Check if access logging is enabled.
Used by protocol handlers to skip building log data when logging is disabled.
"""
return bool(
self.cfg.accesslog or self.cfg.logconfig or
self.cfg.logconfig_dict or self.cfg.logconfig_json or
(self.cfg.syslog and not self.cfg.disable_redirect_access_to_syslog)
)
def access(self, resp, req, environ, request_time):
""" See http://httpd.apache.org/docs/2.0/logs.html#combined
for format details
"""
if not self.access_log_enabled:
return
# wrap atoms:
# - make sure atoms will be test case insensitively
# - if atom doesn't exist replace it by '-'
safe_atoms = self.atoms_wrapper_class(
self.atoms(resp, req, environ, request_time)
)
try:
self.access_log.info(self.cfg.access_log_format, safe_atoms)
except Exception:
self.error(traceback.format_exc())
def now(self):
""" return date in Apache Common Log Format """
return time.strftime('[%d/%b/%Y:%H:%M:%S %z]')
def reopen_files(self):
if self.cfg.capture_output and self.cfg.errorlog != "-":
for stream in sys.stdout, sys.stderr:
stream.flush()
with self.lock:
if self.logfile is not None:
self.logfile.close()
self.logfile = open(self.cfg.errorlog, 'a+')
os.dup2(self.logfile.fileno(), sys.stdout.fileno())
os.dup2(self.logfile.fileno(), sys.stderr.fileno())
for log in loggers():
for handler in log.handlers:
if isinstance(handler, logging.FileHandler):
handler.acquire()
try:
if handler.stream:
handler.close()
handler.stream = handler._open()
finally:
handler.release()
def close_on_exec(self):
for log in loggers():
for handler in log.handlers:
if isinstance(handler, logging.FileHandler):
handler.acquire()
try:
if handler.stream:
util.close_on_exec(handler.stream.fileno())
finally:
handler.release()
def _get_gunicorn_handler(self, log):
for h in log.handlers:
if getattr(h, "_gunicorn", False):
return h
def _set_handler(self, log, output, fmt, stream=None):
# remove previous gunicorn log handler
h = self._get_gunicorn_handler(log)
if h:
log.handlers.remove(h)
if output is not None:
if output == "-":
h = logging.StreamHandler(stream)
else:
util.check_is_writable(output)
h = logging.FileHandler(output)
# make sure the user can reopen the file
try:
os.chown(h.baseFilename, self.cfg.user, self.cfg.group)
except OSError:
# it's probably OK there, we assume the user has given
# /dev/null as a parameter.
pass
h.setFormatter(fmt)
h._gunicorn = True
log.addHandler(h)
def _set_syslog_handler(self, log, cfg, fmt, name):
# setup format
prefix = cfg.syslog_prefix or cfg.proc_name.replace(":", ".")
prefix = "gunicorn.%s.%s" % (prefix, name)
# set format
fmt = logging.Formatter(r"%s: %s" % (prefix, fmt))
# syslog facility
try:
facility = SYSLOG_FACILITIES[cfg.syslog_facility.lower()]
except KeyError as exc:
raise RuntimeError("unknown facility name") from exc
# parse syslog address
socktype, addr = parse_syslog_address(cfg.syslog_addr)
# finally setup the syslog handler
h = logging.handlers.SysLogHandler(address=addr,
facility=facility, socktype=socktype)
h.setFormatter(fmt)
h._gunicorn = True
log.addHandler(h)
def _get_user(self, environ):
user = None
http_auth = environ.get("HTTP_AUTHORIZATION")
if http_auth and http_auth.lower().startswith('basic'):
auth = http_auth.split(" ", 1)
if len(auth) == 2:
try:
# b64decode doesn't accept unicode in Python < 3.3
# so we need to convert it to a byte string
auth = base64.b64decode(auth[1].strip().encode('utf-8'))
# b64decode returns a byte string
user = auth.split(b":", 1)[0].decode("UTF-8")
except (TypeError, binascii.Error, UnicodeDecodeError) as exc:
self.debug("Couldn't get username: %s", exc)
return user
@@ -0,0 +1,36 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
from gunicorn.http.message import Message, Request
from gunicorn.http.parser import RequestParser
def get_parser(cfg, source, source_addr, http2_connection=False):
"""Get appropriate parser based on protocol config.
Args:
cfg: Gunicorn config object
source: Socket or iterable source
source_addr: Source address tuple or None
http2_connection: If True, create HTTP/2 connection handler
Returns:
Parser instance (RequestParser, UWSGIParser, or HTTP2ServerConnection)
"""
# HTTP/2 connection
if http2_connection:
from gunicorn.http2.connection import HTTP2ServerConnection
return HTTP2ServerConnection(cfg, source, source_addr)
# uWSGI protocol
protocol = getattr(cfg, 'protocol', 'http')
if protocol == 'uwsgi':
from gunicorn.uwsgi.parser import UWSGIParser
return UWSGIParser(cfg, source, source_addr)
# Default HTTP/1.x
return RequestParser(cfg, source, source_addr)
__all__ = ['Message', 'Request', 'RequestParser', 'get_parser']
+274
View File
@@ -0,0 +1,274 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import io
import sys
from gunicorn.http.errors import (NoMoreData, ChunkMissingTerminator,
InvalidChunkSize, InvalidChunkExtension)
class ChunkedReader:
def __init__(self, req, unreader):
self.req = req
self.parser = self.parse_chunked(unreader)
self.buf = io.BytesIO()
def read(self, size):
if not isinstance(size, int):
raise TypeError("size must be an integer type")
if size < 0:
raise ValueError("Size must be positive.")
if size == 0:
return b""
if self.parser:
while self.buf.tell() < size:
try:
self.buf.write(next(self.parser))
except StopIteration:
self.parser = None
break
data = self.buf.getvalue()
ret, rest = data[:size], data[size:]
self.buf = io.BytesIO()
self.buf.write(rest)
return ret
def parse_trailers(self, unreader, data):
buf = io.BytesIO()
buf.write(data)
idx = buf.getvalue().find(b"\r\n\r\n")
done = buf.getvalue()[:2] == b"\r\n"
while idx < 0 and not done:
try:
self.get_data(unreader, buf)
except NoMoreData:
# RFC 9112 7.1.2: the last chunk (0 CRLF) must be followed by
# a CRLF-terminated trailer section. Hitting EOF before that
# means the chunked body was truncated, not cleanly ended.
raise ChunkMissingTerminator(b"") from None
idx = buf.getvalue().find(b"\r\n\r\n")
done = buf.getvalue()[:2] == b"\r\n"
if done:
unreader.unread(buf.getvalue()[2:])
return b""
self.req.trailers = self.req.parse_headers(buf.getvalue()[:idx], from_trailer=True)
unreader.unread(buf.getvalue()[idx + 4:])
def parse_chunked(self, unreader):
(size, rest) = self.parse_chunk_size(unreader)
while size > 0:
while size > len(rest):
size -= len(rest)
yield rest
rest = unreader.read()
if not rest:
raise NoMoreData()
yield rest[:size]
# Remove \r\n after chunk
rest = rest[size:]
while len(rest) < 2:
new_data = unreader.read()
if not new_data:
break
rest += new_data
if rest[:2] != b'\r\n':
raise ChunkMissingTerminator(rest[:2])
(size, rest) = self.parse_chunk_size(unreader, data=rest[2:])
def parse_chunk_size(self, unreader, data=None):
buf = io.BytesIO()
if data is not None:
buf.write(data)
idx = buf.getvalue().find(b"\r\n")
while idx < 0:
self.get_data(unreader, buf)
idx = buf.getvalue().find(b"\r\n")
data = buf.getvalue()
line, rest_chunk = data[:idx], data[idx + 2:]
# RFC9112 7.1.1: BWS before chunk-ext - but ONLY then
chunk_size, *chunk_ext = line.split(b";", 1)
if chunk_ext:
# RFC 9112: chunk-ext must not contain bare CR
if b'\r' in chunk_ext[0]:
raise InvalidChunkExtension("bare CR not allowed")
chunk_size = chunk_size.rstrip(b" \t")
if any(n not in b"0123456789abcdefABCDEF" for n in chunk_size):
raise InvalidChunkSize(chunk_size)
if len(chunk_size) == 0:
raise InvalidChunkSize(chunk_size)
chunk_size = int(chunk_size, 16)
if chunk_size == 0:
self.parse_trailers(unreader, rest_chunk)
return (0, None)
return (chunk_size, rest_chunk)
def get_data(self, unreader, buf):
data = unreader.read()
if not data:
raise NoMoreData()
buf.write(data)
class LengthReader:
def __init__(self, unreader, length):
self.unreader = unreader
self.length = length
def read(self, size):
if not isinstance(size, int):
raise TypeError("size must be an integral type")
size = min(self.length, size)
if size < 0:
raise ValueError("Size must be positive.")
if size == 0:
return b""
buf = io.BytesIO()
data = self.unreader.read()
while data:
buf.write(data)
if buf.tell() >= size:
break
data = self.unreader.read()
buf = buf.getvalue()
ret, rest = buf[:size], buf[size:]
self.unreader.unread(rest)
self.length -= size
return ret
class EOFReader:
def __init__(self, unreader):
self.unreader = unreader
self.buf = io.BytesIO()
self.finished = False
def read(self, size):
if not isinstance(size, int):
raise TypeError("size must be an integral type")
if size < 0:
raise ValueError("Size must be positive.")
if size == 0:
return b""
if self.finished:
data = self.buf.getvalue()
ret, rest = data[:size], data[size:]
self.buf = io.BytesIO()
self.buf.write(rest)
return ret
data = self.unreader.read()
while data:
self.buf.write(data)
if self.buf.tell() > size:
break
data = self.unreader.read()
if not data:
self.finished = True
data = self.buf.getvalue()
ret, rest = data[:size], data[size:]
self.buf = io.BytesIO()
self.buf.write(rest)
return ret
class Body:
def __init__(self, reader):
self.reader = reader
self.buf = io.BytesIO()
def __iter__(self):
return self
def __next__(self):
ret = self.readline()
if not ret:
raise StopIteration()
return ret
next = __next__
def getsize(self, size):
if size is None:
return sys.maxsize
elif not isinstance(size, int):
raise TypeError("size must be an integral type")
elif size < 0:
return sys.maxsize
return size
def read(self, size=None):
size = self.getsize(size)
if size == 0:
return b""
if size < self.buf.tell():
data = self.buf.getvalue()
ret, rest = data[:size], data[size:]
self.buf = io.BytesIO()
self.buf.write(rest)
return ret
while size > self.buf.tell():
data = self.reader.read(1024)
if not data:
break
self.buf.write(data)
data = self.buf.getvalue()
ret, rest = data[:size], data[size:]
self.buf = io.BytesIO()
self.buf.write(rest)
return ret
def readline(self, size=None):
size = self.getsize(size)
if size == 0:
return b""
data = self.buf.getvalue()
self.buf = io.BytesIO()
ret = []
while 1:
idx = data.find(b"\n", 0, size)
idx = idx + 1 if idx >= 0 else size if len(data) >= size else 0
if idx:
ret.append(data[:idx])
self.buf.write(data[idx:])
break
ret.append(data)
size -= len(data)
data = self.reader.read(min(1024, size))
if not data:
break
return b"".join(ret)
def readlines(self, size=None):
ret = []
data = self.read()
while data:
pos = data.find(b"\n")
if pos < 0:
ret.append(data)
data = b""
else:
line, data = data[:pos + 1], data[pos + 1:]
ret.append(line)
return ret
+182
View File
@@ -0,0 +1,182 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
# We don't need to call super() in __init__ methods of our
# BaseException and Exception classes because we also define
# our own __str__ methods so there is no need to pass 'message'
# to the base class to get a meaningful output from 'str(exc)'.
# pylint: disable=super-init-not-called
class ParseException(Exception):
pass
class NoMoreData(IOError):
def __init__(self, buf=None):
self.buf = buf
def __str__(self):
return "No more data after: %r" % self.buf
class ConfigurationProblem(ParseException):
def __init__(self, info):
self.info = info
self.code = 500
def __str__(self):
return "Configuration problem: %s" % self.info
class InvalidRequestLine(ParseException):
def __init__(self, req):
self.req = req
self.code = 400
def __str__(self):
return "Invalid HTTP request line: %r" % self.req
class InvalidRequestMethod(ParseException):
def __init__(self, method):
self.method = method
def __str__(self):
return "Invalid HTTP method: %r" % self.method
class ExpectationFailed(ParseException):
def __init__(self, expect):
self.expect = expect
def __str__(self):
return "Unable to comply with expectation: %r" % (self.expect, )
class InvalidHTTPVersion(ParseException):
def __init__(self, version):
self.version = version
def __str__(self):
return "Invalid HTTP Version: %r" % (self.version,)
class InvalidHeader(ParseException):
def __init__(self, hdr, req=None):
self.hdr = hdr
self.req = req
def __str__(self):
return "Invalid HTTP Header: %r" % self.hdr
class ObsoleteFolding(ParseException):
def __init__(self, hdr):
self.hdr = hdr
def __str__(self):
return "Obsolete line folding is unacceptable: %r" % (self.hdr, )
class InvalidHeaderName(ParseException):
def __init__(self, hdr):
self.hdr = hdr
def __str__(self):
return "Invalid HTTP header name: %r" % self.hdr
class UnsupportedTransferCoding(ParseException):
def __init__(self, hdr):
self.hdr = hdr
self.code = 501
def __str__(self):
return "Unsupported transfer coding: %r" % self.hdr
class InvalidChunkSize(IOError):
def __init__(self, data):
self.data = data
def __str__(self):
return "Invalid chunk size: %r" % self.data
class ChunkMissingTerminator(IOError):
def __init__(self, term):
self.term = term
def __str__(self):
return "Invalid chunk terminator is not '\\r\\n': %r" % self.term
class InvalidChunkExtension(IOError):
"""Invalid chunk extension per RFC 9112."""
def __init__(self, reason):
self.reason = reason
def __str__(self):
return "Invalid chunk extension: %s" % self.reason
class LimitRequestLine(ParseException):
def __init__(self, size, max_size=None):
self.size = size
self.max_size = max_size
def __str__(self):
if self.max_size is None:
return str(self.size)
return "Request Line is too large (%s > %s)" % (self.size, self.max_size)
class LimitRequestHeaders(ParseException):
def __init__(self, msg):
self.msg = msg
def __str__(self):
return self.msg
class InvalidProxyLine(ParseException):
def __init__(self, line):
self.line = line
self.code = 400
def __str__(self):
return "Invalid PROXY line: %r" % self.line
class InvalidProxyHeader(ParseException):
def __init__(self, msg):
self.msg = msg
self.code = 400
def __str__(self):
return "Invalid PROXY header: %s" % self.msg
class ForbiddenProxyRequest(ParseException):
def __init__(self, host):
self.host = host
self.code = 403
def __str__(self):
return "Proxy request from %r not allowed" % self.host
class InvalidSchemeHeaders(ParseException):
def __str__(self):
return "Contradictory scheme headers"
class InvalidH2CPreface(ParseException):
def __init__(self, data):
self.data = data
def __str__(self):
return "Expected HTTP/2 connection preface, got %r" % (self.data,)
+939
View File
@@ -0,0 +1,939 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
from enum import IntEnum
import ipaddress
import re
import socket
import struct
from gunicorn.http.body import ChunkedReader, LengthReader, EOFReader, Body
from gunicorn.http.errors import (
InvalidHeader, InvalidHeaderName, NoMoreData,
InvalidRequestLine, InvalidRequestMethod, InvalidHTTPVersion,
LimitRequestLine, LimitRequestHeaders,
UnsupportedTransferCoding, ObsoleteFolding,
ExpectationFailed,
)
from gunicorn.http.errors import InvalidProxyLine, InvalidProxyHeader, ForbiddenProxyRequest
from gunicorn.http.errors import InvalidSchemeHeaders
from gunicorn.util import bytes_to_str, split_request_uri
# Fast parser availability (cached at module level)
_fast_parser_available = None
_fast_parser_module = None
# Compatibility flags not supported by the fast parser
_FAST_PARSER_INCOMPATIBLE_FLAGS = (
'permit_obsolete_folding',
'strip_header_spaces',
)
def _check_fast_parser(cfg):
"""Check if fast C parser is available and should be used.
Returns False if:
- http_parser='python' is explicitly set
- gunicorn_h1c is not installed (in 'auto' mode)
- gunicorn_h1c < 0.4.1 (in 'auto' mode)
- Incompatible compatibility flags are enabled (in 'auto' mode)
Raises RuntimeError if:
- http_parser='fast' but gunicorn_h1c is not installed
- http_parser='fast' but gunicorn_h1c < 0.4.1
- http_parser='fast' but incompatible flags are enabled
"""
global _fast_parser_available, _fast_parser_module # pylint: disable=global-statement
parser_setting = getattr(cfg, 'http_parser', 'auto')
if parser_setting == 'python':
return False
if _fast_parser_available is None:
try:
import gunicorn_h1c
_fast_parser_available = True
_fast_parser_module = gunicorn_h1c
except ImportError:
_fast_parser_available = False
if not _fast_parser_available and parser_setting == 'fast':
raise RuntimeError("gunicorn_h1c not installed but http_parser='fast'")
if not _fast_parser_available:
return False
# Require >= 0.4.1 for limit enforcement
if not hasattr(_fast_parser_module, 'LimitRequestLine'):
if parser_setting == 'fast':
raise RuntimeError(
"gunicorn_h1c >= 0.4.1 required for http_parser='fast'. "
"Please upgrade: pip install --upgrade gunicorn_h1c"
)
# In 'auto' mode, fall back to Python parser
return False
# Check for incompatible compatibility flags
incompatible = []
for flag in _FAST_PARSER_INCOMPATIBLE_FLAGS:
if getattr(cfg, flag, False):
incompatible.append(flag)
if incompatible:
if parser_setting == 'fast':
raise RuntimeError(
"http_parser='fast' is incompatible with compatibility flags: %s. "
"Use http_parser='python' or disable these flags."
% ', '.join(incompatible)
)
# In 'auto' mode, fall back to Python parser
return False
return True
# PROXY protocol v2 constants
PP_V2_SIGNATURE = b"\x0D\x0A\x0D\x0A\x00\x0D\x0A\x51\x55\x49\x54\x0A"
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
MAX_REQUEST_LINE = 8190
MAX_HEADERS = 32768
DEFAULT_MAX_HEADERFIELD_SIZE = 8190
# verbosely on purpose, avoid backslash ambiguity
RFC9110_5_6_2_TOKEN_SPECIALS = r"!#$%&'*+-.^_`|~"
TOKEN_RE = re.compile(r"[%s0-9a-zA-Z]+" % (re.escape(RFC9110_5_6_2_TOKEN_SPECIALS)))
METHOD_BADCHAR_RE = re.compile("[a-z#]")
# usually 1.0 or 1.1 - RFC9112 permits restricting to single-digit versions
VERSION_RE = re.compile(r"HTTP/(\d)\.(\d)")
# RFC 9110 section 5.5: field-vchar = VCHAR / obs-text; SP and HTAB are the
# only non-VCHAR bytes allowed in a field-value. Anything else in the
# control range (0x00-0x1F except HTAB, plus DEL 0x7F) must be rejected.
RFC9110_5_5_INVALID_AND_DANGEROUS = re.compile(r"[\x00-\x08\x0a-\x1f\x7f]")
# RFC 9110 section 6.5.1: fields forbidden in trailers because they alter
# routing, framing, or authentication. Using the uppercased names stored
# by parse_headers.
RFC9110_6_5_1_FORBIDDEN_TRAILER = frozenset((
"HOST",
"CONTENT-LENGTH",
"TRANSFER-ENCODING",
"TRAILER",
"AUTHORIZATION",
"TE",
))
# 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.
# CONTENT-LENGTH is deliberately absent: duplicates are already rejected in
# set_body_reader(), which reports them with the request attached.
RFC9110_5_3_SINGLETON_FIELDS = frozenset((
"HOST",
"CONTENT-TYPE",
))
def _ip_in_allow_list(ip_str, allow_list, networks):
"""Check if IP address is in the allow list.
Args:
ip_str: The IP address string to check
allow_list: The original allow list (strings, may contain "*")
networks: Pre-computed ipaddress.ip_network objects from config
"""
if '*' in allow_list:
return True
try:
ip = ipaddress.ip_address(ip_str)
except ValueError:
return False
for network in networks:
if ip in network:
return True
return False
class HeaderPolicy:
"""Header policy shared by every request path.
Applies to HTTP/1 (both parsers) and HTTP/2 alike, so a rule added here
cannot be enforced on one protocol and quietly skipped on another.
Requires ``self.cfg``, ``self.peer_addr``, ``self.scheme`` and
``self.version`` to be set before any header is applied.
"""
#: Set by whichever class mixes this in, before headers are applied.
scheme = None
version = None
_expected_100_continue = False
#: HTTP/2 has no 100-continue handshake on the wire the way HTTP/1 does,
#: and gunicorn answers one by writing HTTP/1 bytes straight to the socket
#: (see wsgi.create), which would corrupt an HTTP/2 connection.
_policy_expect_continue = True
def _peer_is_trusted_proxy(self):
"""Whether the peer may set forwarding and scheme headers."""
cfg = self.cfg
return (not isinstance(self.peer_addr, tuple)
or _ip_in_allow_list(self.peer_addr[0], cfg.forwarded_allow_ips,
cfg.forwarded_allow_networks()))
def _peer_trusted_for_forwarded(self):
"""Return the (secure_scheme_headers, forwarder_headers) the peer is allowed to set.
When the peer's address is not in ``forwarded_allow_ips`` (or networks),
configured forwarding/secure-scheme policy must be ignored to prevent
spoofing. Returns ``({}, [])`` when the peer is untrusted.
"""
cfg = self.cfg
if self._peer_is_trusted_proxy():
return cfg.secure_scheme_headers, cfg.forwarder_headers
return {}, []
def _apply_header_policy(self, name, value, scheme_state, seen,
secure_scheme_headers, forwarder_headers,
from_trailer=False):
"""Apply per-header policy shared between Python and fast parsers.
Mutates ``self._expected_100_continue`` and ``self.scheme`` as needed.
``scheme_state`` is a single-element list used as a mutable sentinel
so the caller can detect repeated scheme headers. ``seen`` is a set of
the field names already accepted for this message, used the same way to
detect repeated singleton fields.
Returns the (name, value) pair to retain, or ``None`` to drop the
header (per ``header_map='drop'``). Raises the same exceptions the
Python path raises so behavior is identical regardless of parser.
"""
# https://datatracker.ietf.org/doc/html/rfc9110#section-5.3
# A singleton field cannot be combined into a list, so a repeat leaves
# the message ambiguous and open to being read differently by us and by
# anything downstream.
if name in RFC9110_5_3_SINGLETON_FIELDS:
if name in seen:
raise InvalidHeader(name, req=self)
seen.add(name)
if (self._policy_expect_continue and not from_trailer
and name == "EXPECT"):
# https://datatracker.ietf.org/doc/html/rfc9110#section-10.1.1
# "The Expect field value is case-insensitive."
if value.lower() == "100-continue":
if self.version < (1, 1):
# https://datatracker.ietf.org/doc/html/rfc9110#section-10.1.1-12
# "A server that receives a 100-continue expectation
# in an HTTP/1.0 request MUST ignore that expectation."
pass
else:
self._expected_100_continue = True
# N.B. understood but ignored expect header does not return 417
else:
raise ExpectationFailed(value)
if name in secure_scheme_headers:
secure = value == secure_scheme_headers[name]
scheme = "https" if secure else "http"
if scheme_state[0]:
if scheme != self.scheme:
raise InvalidSchemeHeaders()
else:
scheme_state[0] = True
self.scheme = scheme
# ambiguous mapping allows fooling downstream, e.g. merging non-identical headers:
# X-Forwarded-For: 2001:db8::ha:cc:ed
# X_Forwarded_For: 127.0.0.1,::1
# HTTP_X_FORWARDED_FOR = 2001:db8::ha:cc:ed,127.0.0.1,::1
# Only modify after fixing *ALL* header transformations; network to wsgi env
if "_" in name:
if name in forwarder_headers or "*" in forwarder_headers:
# This forwarder may override our environment
pass
elif self.cfg.header_map == "dangerous":
# as if we did not know we cannot safely map this
pass
elif self.cfg.header_map == "drop":
# almost as if it never had been there
# but still counts against resource limits
return None
else:
# fail-safe fallthrough: refuse
raise InvalidHeaderName(name)
return (name, value)
class Message(HeaderPolicy):
def __init__(self, cfg, unreader, peer_addr):
self.cfg = cfg
self.unreader = unreader
self.peer_addr = peer_addr
self.remote_addr = peer_addr
self.version = None
self.headers = []
self.trailers = []
self.body = None
self.scheme = "https" if cfg.is_ssl else "http"
self.must_close = False
self._expected_100_continue = False
# set headers limits
self.limit_request_fields = cfg.limit_request_fields
if (self.limit_request_fields <= 0
or self.limit_request_fields > MAX_HEADERS):
self.limit_request_fields = MAX_HEADERS
self.limit_request_field_size = cfg.limit_request_field_size
if self.limit_request_field_size <= 0:
self.limit_request_field_size = DEFAULT_MAX_HEADERFIELD_SIZE
# set max header buffer size
max_header_field_size = self.limit_request_field_size or DEFAULT_MAX_HEADERFIELD_SIZE
self.max_buffer_headers = self.limit_request_fields * \
(max_header_field_size + 2) + 4
unused = self.parse(self.unreader)
self.unreader.unread(unused)
self.set_body_reader()
def force_close(self):
self.must_close = True
def parse(self, unreader):
raise NotImplementedError()
def parse_headers(self, data, from_trailer=False):
headers = []
# Split lines on \r\n
lines = [bytes_to_str(line) for line in data.split(b"\r\n")]
# handle scheme headers
scheme_state = [False]
seen = set()
if from_trailer:
# nonsense. either a request is https from the beginning
# .. or we are just behind a proxy who does not remove conflicting trailers
secure_scheme_headers, forwarder_headers = {}, []
else:
secure_scheme_headers, forwarder_headers = self._peer_trusted_for_forwarded()
# Parse headers into key/value pairs paying attention
# to continuation lines.
while lines:
if len(headers) >= self.limit_request_fields:
raise LimitRequestHeaders("limit request headers fields")
# Parse initial header name: value pair.
curr = lines.pop(0)
header_length = len(curr) + len("\r\n")
if curr.find(":") <= 0:
raise InvalidHeader(curr)
name, value = curr.split(":", 1)
if self.cfg.strip_header_spaces:
name = name.rstrip(" \t")
if not TOKEN_RE.fullmatch(name):
raise InvalidHeaderName(name)
# this is still a dangerous place to do this
# but it is more correct than doing it before the pattern match:
# after we entered Unicode wonderland, 8bits could case-shift into ASCII:
# b"\xDF".decode("latin-1").upper().encode("ascii") == b"SS"
name = name.upper()
# RFC 9110 section 6.5.1
if from_trailer and name in RFC9110_6_5_1_FORBIDDEN_TRAILER:
raise InvalidHeaderName(name)
value = [value.strip(" \t")]
# Consume value continuation lines..
while lines and lines[0].startswith((" ", "\t")):
# .. which is obsolete here, and no longer done by default
if not self.cfg.permit_obsolete_folding:
raise ObsoleteFolding(name)
curr = lines.pop(0)
header_length += len(curr) + len("\r\n")
if header_length > self.limit_request_field_size > 0:
raise LimitRequestHeaders("limit request headers "
"fields size")
value.append(curr.strip("\t "))
value = " ".join(value)
if RFC9110_5_5_INVALID_AND_DANGEROUS.search(value):
raise InvalidHeader(name)
if header_length > self.limit_request_field_size > 0:
raise LimitRequestHeaders("limit request headers fields size")
kept = self._apply_header_policy(
name, value, scheme_state, seen,
secure_scheme_headers, forwarder_headers,
from_trailer=from_trailer,
)
if kept is None:
continue
headers.append(kept)
return headers
def set_body_reader(self):
chunked = False
content_length = None
for (name, value) in self.headers:
if name == "CONTENT-LENGTH":
if content_length is not None:
raise InvalidHeader("CONTENT-LENGTH", req=self)
content_length = value
elif name == "TRANSFER-ENCODING":
# T-E can be a list
# https://datatracker.ietf.org/doc/html/rfc9112#name-transfer-encoding
vals = [v.strip() for v in value.split(',')]
for val in vals:
if val.lower() == "chunked":
# DANGER: transfer codings stack, and stacked chunking is never intended
if chunked:
raise InvalidHeader("TRANSFER-ENCODING", req=self)
chunked = True
elif val.lower() == "identity":
# does not do much, could still plausibly desync from what the proxy does
# safe option: nuke it, its never needed
if chunked:
raise InvalidHeader("TRANSFER-ENCODING", req=self)
elif val.lower() in ('compress', 'deflate', 'gzip'):
# chunked should be the last one
if chunked:
raise InvalidHeader("TRANSFER-ENCODING", req=self)
self.force_close()
else:
raise UnsupportedTransferCoding(value)
if chunked:
# two potentially dangerous cases:
# a) CL + TE (TE overrides CL.. only safe if the recipient sees it that way too)
# b) chunked HTTP/1.0 (always faulty)
if self.version < (1, 1):
# framing wonky, see RFC 9112 Section 6.1
raise InvalidHeader("TRANSFER-ENCODING", req=self)
if content_length is not None:
# we cannot be certain the message framing we understood matches proxy intent
# -> whatever happens next, remaining input must not be trusted
raise InvalidHeader("CONTENT-LENGTH", req=self)
self.body = Body(ChunkedReader(self, self.unreader))
elif content_length is not None:
try:
if str(content_length).isnumeric():
content_length = int(content_length)
else:
raise InvalidHeader("CONTENT-LENGTH", req=self)
except ValueError:
raise InvalidHeader("CONTENT-LENGTH", req=self)
if content_length < 0:
raise InvalidHeader("CONTENT-LENGTH", req=self)
self.body = Body(LengthReader(self.unreader, content_length))
else:
self.body = Body(EOFReader(self.unreader))
def should_close(self):
if self.must_close:
return True
for (h, v) in self.headers:
if h == "CONNECTION":
v = v.lower().strip(" \t")
if v == "close":
return True
elif v == "keep-alive":
return False
break
return self.version <= (1, 0)
class Request(Message):
def __init__(self, cfg, unreader, peer_addr, req_number=1):
self.method = None
self.uri = None
self.path = None
self.query = None
self.fragment = None
# get max request line size (0 means unlimited per documentation)
self.limit_request_line = cfg.limit_request_line
if self.limit_request_line < 0:
self.limit_request_line = MAX_REQUEST_LINE
# For fast parser: use large value when unlimited (0), since C parser
# doesn't support 0 as unlimited. 1MB should be more than enough.
if self.limit_request_line == 0:
self._fast_limit_request_line = 1024 * 1024 # 1MB
elif self.limit_request_line >= MAX_REQUEST_LINE:
self._fast_limit_request_line = MAX_REQUEST_LINE
self.limit_request_line = MAX_REQUEST_LINE
else:
self._fast_limit_request_line = self.limit_request_line
self.req_number = req_number
self.proxy_protocol_info = None
# Check if fast parser should be used
self._use_fast = _check_fast_parser(cfg)
super().__init__(cfg, unreader, peer_addr)
def get_data(self, unreader, buf, stop=False):
data = unreader.read()
if not data:
if stop:
raise StopIteration()
raise NoMoreData(buf.getvalue())
buf.write(data)
def parse(self, unreader):
buf = bytearray()
self.read_into(unreader, buf, stop=True)
# Handle proxy protocol if enabled and this is the first request
mode = self.cfg.proxy_protocol
if mode != "off" and self.req_number == 1:
buf = self._handle_proxy_protocol(unreader, buf, mode)
# Use fast parser if available
if self._use_fast:
return self._parse_fast(unreader, buf)
return self._parse_python(unreader, buf)
def _parse_fast(self, unreader, buf):
"""Parse request using fast C parser (gunicorn_h1c >= 0.4.1)."""
# Read until we have complete headers
data = bytes(buf)
last_len = 0
while True:
try:
# Pass all limit parameters to C parser
# Use _fast_limit_request_line which handles 0=unlimited
result = _fast_parser_module.parse_request(
data,
last_len=last_len,
limit_request_line=self._fast_limit_request_line,
limit_request_fields=self.limit_request_fields,
limit_request_field_size=self.limit_request_field_size,
permit_unconventional_http_method=self.cfg.permit_unconventional_http_method,
permit_unconventional_http_version=self.cfg.permit_unconventional_http_version,
)
break
except _fast_parser_module.IncompleteError:
last_len = len(data)
self.read_into(unreader, buf)
data = bytes(buf)
if len(data) > self.max_buffer_headers + self._fast_limit_request_line:
raise LimitRequestHeaders("max buffer headers")
except _fast_parser_module.LimitRequestLine as e:
raise LimitRequestLine(str(e))
except _fast_parser_module.LimitRequestHeaders as e:
raise LimitRequestHeaders(str(e))
except _fast_parser_module.InvalidRequestMethod as e:
raise InvalidRequestMethod(str(e))
except _fast_parser_module.InvalidHTTPVersion as e:
raise InvalidHTTPVersion(str(e))
except _fast_parser_module.InvalidHeaderName as e:
raise InvalidHeaderName(str(e))
except _fast_parser_module.InvalidHeader as e:
raise InvalidHeader(str(e))
except _fast_parser_module.ParseError as e:
raise InvalidRequestLine(str(e))
# Extract parsed data
self.method = bytes_to_str(result['method'])
self.uri = bytes_to_str(result['path'])
# Casefold method if configured (validation done by C parser)
if self.cfg.casefold_http_method:
self.method = self.method.upper()
# Parse URI parts
if len(self.uri) == 0:
raise InvalidRequestLine(self.uri)
try:
parts = split_request_uri(self.uri)
except ValueError:
raise InvalidRequestLine(self.uri)
self.path = parts.path or ""
self.query = parts.query or ""
self.fragment = parts.fragment or ""
# Version (validation done by C parser)
self.version = (1, result['minor_version'])
# Headers - convert bytes to strings with uppercase names
# gunicorn_h1c returns headers as (bytes, bytes) tuples
# Header name/value validation done by C parser; policy (Expect,
# secure_scheme_headers, forwarder trust gate, header_map) is enforced
# below so the fast path mirrors parse_headers().
self.headers = []
scheme_state = [False]
seen = set()
secure_scheme_headers, forwarder_headers = self._peer_trusted_for_forwarded()
for name_bytes, value_bytes in result['headers']:
name = bytes_to_str(name_bytes).upper()
value = bytes_to_str(value_bytes)
kept = self._apply_header_policy(
name, value, scheme_state, seen,
secure_scheme_headers, forwarder_headers,
)
if kept is None:
continue
self.headers.append(kept)
# Return remaining data after headers
consumed = result['consumed']
return data[consumed:]
def _parse_python(self, unreader, buf):
"""Parse request using pure Python parser."""
# Get request line
line, buf = self.read_line(unreader, buf, self.limit_request_line)
self.parse_request_line(line)
# Headers
data = bytes(buf)
done = data[:2] == b"\r\n"
while True:
idx = data.find(b"\r\n\r\n")
done = data[:2] == b"\r\n"
if idx < 0 and not done:
self.read_into(unreader, buf)
data = bytes(buf)
if len(data) > self.max_buffer_headers:
raise LimitRequestHeaders("max buffer headers")
else:
break
if done:
self.unreader.unread(data[2:])
return b""
self.headers = self.parse_headers(data[:idx], from_trailer=False)
ret = data[idx + 4:]
return ret
def read_into(self, unreader, buf, stop=False):
"""Read data from unreader and append to bytearray buffer."""
data = unreader.read()
if not data:
if stop:
raise StopIteration()
raise NoMoreData(bytes(buf))
buf.extend(data)
def read_line(self, unreader, buf, limit=0):
"""Read a line from buffer, returning (line, remaining_buffer)."""
data = bytes(buf)
while True:
idx = data.find(b"\r\n")
if idx >= 0:
# check if the request line is too large
if idx > limit > 0:
raise LimitRequestLine(idx, limit)
break
if len(data) - 2 > limit > 0:
raise LimitRequestLine(len(data), limit)
self.read_into(unreader, buf)
data = bytes(buf)
return (data[:idx], # request line,
bytearray(data[idx + 2:])) # residue in the buffer, skip \r\n
def read_bytes(self, unreader, buf, count):
"""Read exactly count bytes from buffer/unreader."""
while len(buf) < count:
self.read_into(unreader, buf)
return bytes(buf[:count]), bytearray(buf[count:])
def _handle_proxy_protocol(self, unreader, buf, mode):
"""Handle PROXY protocol detection and parsing.
Returns the buffer with proxy protocol data consumed.
"""
# Ensure we have enough data to detect v2 signature (12 bytes)
while len(buf) < 12:
self.read_into(unreader, buf)
# Check for v2 signature first
if mode in ("v2", "auto") and buf[:12] == PP_V2_SIGNATURE:
self.proxy_protocol_access_check()
return self._parse_proxy_protocol_v2(unreader, buf)
# Check for v1 prefix
if mode in ("v1", "auto") and buf[:6] == b"PROXY ":
self.proxy_protocol_access_check()
return self._parse_proxy_protocol_v1(unreader, buf)
# Not proxy protocol - return buffer unchanged
return buf
def proxy_protocol_access_check(self):
"""Check if proxy protocol is allowed from this peer."""
if (isinstance(self.peer_addr, tuple) and
not _ip_in_allow_list(self.peer_addr[0], self.cfg.proxy_allow_ips,
self.cfg.proxy_allow_networks())):
raise ForbiddenProxyRequest(self.peer_addr[0])
def _parse_proxy_protocol_v1(self, unreader, buf):
"""Parse PROXY protocol v1 (text format).
Returns buffer with v1 header consumed.
"""
# Read until we find \r\n
data = bytes(buf)
while b"\r\n" not in data:
self.read_into(unreader, buf)
data = bytes(buf)
idx = data.find(b"\r\n")
line = bytes_to_str(data[:idx])
remaining = bytearray(data[idx + 2:])
bits = line.split(" ")
if len(bits) != 6:
raise InvalidProxyLine(line)
# Extract data
proto = bits[1]
s_addr = bits[2]
d_addr = bits[3]
# Validation
if proto not in ["TCP4", "TCP6"]:
raise InvalidProxyLine("protocol '%s' not supported" % proto)
if proto == "TCP4":
try:
socket.inet_pton(socket.AF_INET, s_addr)
socket.inet_pton(socket.AF_INET, d_addr)
except OSError:
raise InvalidProxyLine(line)
elif proto == "TCP6":
try:
socket.inet_pton(socket.AF_INET6, s_addr)
socket.inet_pton(socket.AF_INET6, d_addr)
except OSError:
raise InvalidProxyLine(line)
try:
s_port = int(bits[4])
d_port = int(bits[5])
except ValueError:
raise InvalidProxyLine("invalid port %s" % line)
if not ((0 <= s_port <= 65535) and (0 <= d_port <= 65535)):
raise InvalidProxyLine("invalid port %s" % line)
# Set data
self.proxy_protocol_info = {
"proxy_protocol": proto,
"client_addr": s_addr,
"client_port": s_port,
"proxy_addr": d_addr,
"proxy_port": d_port
}
return remaining
def _parse_proxy_protocol_v2(self, unreader, buf):
"""Parse PROXY protocol v2 (binary format).
Returns buffer with v2 header consumed.
"""
# We need at least 16 bytes for the header (12 signature + 4 header)
while len(buf) < 16:
self.read_into(unreader, buf)
# Parse header fields (after 12-byte signature)
ver_cmd = buf[12]
fam_proto = buf[13]
length = struct.unpack(">H", bytes(buf[14:16]))[0]
# Validate version (high nibble must be 0x2)
version = (ver_cmd & 0xF0) >> 4
if version != 2:
raise InvalidProxyHeader("unsupported version %d" % version)
# Extract command (low nibble)
command = ver_cmd & 0x0F
if command not in (PPCommand.LOCAL, PPCommand.PROXY):
raise InvalidProxyHeader("unsupported command %d" % command)
# Ensure we have the complete header
total_header_size = 16 + length
while len(buf) < total_header_size:
self.read_into(unreader, buf)
# For LOCAL command, no address info is provided
if command == PPCommand.LOCAL:
self.proxy_protocol_info = {
"proxy_protocol": "LOCAL",
"client_addr": None,
"client_port": None,
"proxy_addr": None,
"proxy_port": None
}
return bytearray(buf[total_header_size:])
# Extract address family and protocol
family = (fam_proto & 0xF0) >> 4
protocol = fam_proto & 0x0F
# We only support TCP (STREAM)
if protocol != PPProtocol.STREAM:
raise InvalidProxyHeader("only TCP protocol is supported")
addr_data = bytes(buf[16:16 + length])
if family == PPFamily.INET: # IPv4
if length < 12: # 4+4+2+2
raise InvalidProxyHeader("insufficient address data for IPv4")
s_addr = socket.inet_ntop(socket.AF_INET, addr_data[0:4])
d_addr = socket.inet_ntop(socket.AF_INET, 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 length < 36: # 16+16+2+2
raise InvalidProxyHeader("insufficient address data for IPv6")
s_addr = socket.inet_ntop(socket.AF_INET6, addr_data[0:16])
d_addr = socket.inet_ntop(socket.AF_INET6, addr_data[16:32])
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:
# No address info provided with PROXY command
self.proxy_protocol_info = {
"proxy_protocol": "UNSPEC",
"client_addr": None,
"client_port": None,
"proxy_addr": None,
"proxy_port": None
}
return bytearray(buf[total_header_size:])
else:
raise InvalidProxyHeader("unsupported address family %d" % family)
# Set data
self.proxy_protocol_info = {
"proxy_protocol": proto,
"client_addr": s_addr,
"client_port": s_port,
"proxy_addr": d_addr,
"proxy_port": d_port
}
return bytearray(buf[total_header_size:])
def parse_request_line(self, line_bytes):
bits = [bytes_to_str(bit) for bit in line_bytes.split(b" ", 2)]
if len(bits) != 3:
raise InvalidRequestLine(bytes_to_str(line_bytes))
# Method: RFC9110 Section 9
self.method = bits[0]
# nonstandard restriction, suitable for all IANA registered methods
# partially enforced in previous gunicorn versions
if not self.cfg.permit_unconventional_http_method:
if METHOD_BADCHAR_RE.search(self.method):
raise InvalidRequestMethod(self.method)
if not 3 <= len(bits[0]) <= 20:
raise InvalidRequestMethod(self.method)
# standard restriction: RFC9110 token
if not TOKEN_RE.fullmatch(self.method):
raise InvalidRequestMethod(self.method)
# nonstandard and dangerous
# methods are merely uppercase by convention, no case-insensitive treatment is intended
if self.cfg.casefold_http_method:
self.method = self.method.upper()
# URI
self.uri = bits[1]
# Python stdlib explicitly tells us it will not perform validation.
# https://docs.python.org/3/library/urllib.parse.html#url-parsing-security
# There are *four* `request-target` forms in rfc9112, none of them can be empty:
# 1. origin-form, which starts with a slash
# 2. absolute-form, which starts with a non-empty scheme
# 3. authority-form, (for CONNECT) which contains a colon after the host
# 4. asterisk-form, which is an asterisk (`\x2A`)
# => manually reject one always invalid URI: empty
if len(self.uri) == 0:
raise InvalidRequestLine(bytes_to_str(line_bytes))
# RFC 9112 section 3.2.4: asterisk-form is only valid with OPTIONS.
if self.uri == "*" and self.method != "OPTIONS":
raise InvalidRequestLine(bytes_to_str(line_bytes))
# RFC 9112 section 3.2.3: authority-form ("host:port") is only valid
# with CONNECT. origin-form starts with "/"; absolute-form contains
# "://". Anything else on a non-CONNECT request is authority-form.
if (self.method != "CONNECT"
and self.uri != "*"
and not self.uri.startswith("/")
and "://" not in self.uri):
raise InvalidRequestLine(bytes_to_str(line_bytes))
try:
parts = split_request_uri(self.uri)
except ValueError:
raise InvalidRequestLine(bytes_to_str(line_bytes))
self.path = parts.path or ""
self.query = parts.query or ""
self.fragment = parts.fragment or ""
# Version
match = VERSION_RE.fullmatch(bits[2])
if match is None:
raise InvalidHTTPVersion(bits[2])
self.version = (int(match.group(1)), int(match.group(2)))
if not (1, 0) <= self.version < (2, 0):
# if ever relaxing this, carefully review Content-Encoding processing
if not self.cfg.permit_unconventional_http_version:
raise InvalidHTTPVersion(self.version)
def set_body_reader(self):
super().set_body_reader()
if isinstance(self.body.reader, EOFReader):
self.body = Body(LengthReader(self.unreader, 0))
+121
View File
@@ -0,0 +1,121 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import socket
import ssl
import time
from gunicorn.http.message import Request
from gunicorn.http.unreader import SocketUnreader, IterUnreader
# Cap on bytes drained from an unconsumed request body before a keepalive
# reset. Defends against a slow-but-steady client that stays under a per-read
# deadline yet streams indefinitely.
_DRAIN_MAX_BYTES = 64 * 1024
class Parser:
mesg_class = None
def __init__(self, cfg, source, source_addr):
self.cfg = cfg
if hasattr(source, "recv"):
self.unreader = SocketUnreader(source)
else:
self.unreader = IterUnreader(source)
self.mesg = None
self.source_addr = source_addr
# request counter (for keepalive connetions)
self.req_count = 0
def __iter__(self):
return self
def finish_body(self, deadline=None, max_bytes=None):
"""Discard any unread body of the current message.
Called before returning a keepalive connection to the poller so the
socket does not appear readable due to leftover body bytes.
``deadline`` is an absolute ``time.monotonic()`` value; when set the
socket read timeout is bounded by the remaining time before each read.
``max_bytes`` caps the total drained bytes; when a deadline is given
and ``max_bytes`` is left at the default, ``_DRAIN_MAX_BYTES`` applies
to defend against a slow client that keeps trickling under it. When
called without a deadline (the default invocation from ``__next__``),
no byte cap is applied so the prior unbounded drain semantics are
preserved for callers that don't know how to react to a partial drain.
Returns ``True`` when the body was fully drained, ``False`` when the
drain was abandoned (deadline, byte cap, or socket timeout). Callers
that observe ``False`` MUST close the connection rather than serve
another request on it.
"""
if not self.mesg:
return True
if max_bytes is None and deadline is not None:
max_bytes = _DRAIN_MAX_BYTES
sock = getattr(self.unreader, "sock", None)
# gettimeout/settimeout only matter when bounding a real socket; a
# mock or non-socket source skips the timeout plumbing.
if sock is not None and hasattr(sock, "gettimeout") and hasattr(sock, "settimeout"):
timeoutable_sock = sock
prior_timeout = sock.gettimeout()
else:
timeoutable_sock = None
prior_timeout = None
drained = 0
try:
while True:
if deadline is not None and timeoutable_sock is not None:
remaining = deadline - time.monotonic()
if remaining <= 0:
return False
timeoutable_sock.settimeout(remaining)
try:
data = self.mesg.body.read(1024)
except (socket.timeout, TimeoutError):
return False
except ssl.SSLWantReadError:
# SSL socket has no more application data available
return True
if not data:
return True
drained += len(data)
if max_bytes is not None and drained >= max_bytes:
return False
finally:
if timeoutable_sock is not None:
try:
timeoutable_sock.settimeout(prior_timeout)
except OSError:
pass
def __next__(self):
# Stop if HTTP dictates a stop.
if self.mesg and self.mesg.should_close():
raise StopIteration()
# Discard any unread body of the previous message
self.finish_body()
# Parse the next request
self.req_count += 1
self.mesg = self.mesg_class(self.cfg, self.unreader, self.source_addr, self.req_count)
if not self.mesg:
raise StopIteration()
return self.mesg
next = __next__
class RequestParser(Parser):
mesg_class = Request
@@ -0,0 +1,92 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import io
import os
# Classes that can undo reading data from
# a given type of data source.
class Unreader:
def __init__(self):
self.buf = io.BytesIO()
def chunk(self):
raise NotImplementedError()
def take_buffered(self):
"""Return read-ahead already held, without touching the source.
read() blocks on the source when the buffer is empty, which is wrong
for a caller that only wants the bytes it has: an Upgrade: h2c
handshake needs whatever the client pipelined behind the request,
and must not wait for more.
"""
data = self.buf.getvalue()
self.buf = io.BytesIO()
return data
def read(self, size=None):
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
self.buf.seek(0, os.SEEK_END)
if size is None and self.buf.tell():
ret = self.buf.getvalue()
self.buf = io.BytesIO()
return ret
if size is None:
d = self.chunk()
return d
while self.buf.tell() < size:
chunk = self.chunk()
if not chunk:
ret = self.buf.getvalue()
self.buf = io.BytesIO()
return ret
self.buf.write(chunk)
data = self.buf.getvalue()
self.buf = io.BytesIO()
self.buf.write(data[size:])
return data[:size]
def unread(self, data):
rest = self.buf.getvalue()
self.buf = io.BytesIO()
self.buf.write(data)
self.buf.write(rest)
class SocketUnreader(Unreader):
def __init__(self, sock, max_chunk=8192):
super().__init__()
self.sock = sock
self.mxchunk = max_chunk
def chunk(self):
return self.sock.recv(self.mxchunk)
class IterUnreader(Unreader):
def __init__(self, iterable):
super().__init__()
self.iter = iter(iterable)
def chunk(self):
if not self.iter:
return b""
try:
return next(self.iter)
except StopIteration:
self.iter = None
return b""
+551
View File
@@ -0,0 +1,551 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
import io
import logging
import os
import re
import sys
from gunicorn.http.message import TOKEN_RE
from gunicorn.http.errors import ConfigurationProblem, InvalidHeader, InvalidHeaderName
from gunicorn import SERVER_SOFTWARE, SERVER
from gunicorn import util
# Send files in at most 1GB blocks as some operating systems can have problems
# with sending files in blocks over 2GB.
BLKSIZE = 0x3FFFFFFF
# RFC9110 5.5: field-vchar = VCHAR / obs-text
# RFC4234 B.1: VCHAR = 0x21-x07E = printable ASCII
HEADER_VALUE_RE = re.compile(r'[ \t\x21-\x7e\x80-\xff]*')
log = logging.getLogger(__name__)
class FileWrapper:
def __init__(self, filelike, blksize=8192):
self.filelike = filelike
self.blksize = blksize
if hasattr(filelike, 'close'):
self.close = filelike.close
def __getitem__(self, key):
data = self.filelike.read(self.blksize)
if data:
return data
raise IndexError
def __iter__(self):
return self
def __next__(self):
data = self.filelike.read(self.blksize)
if data:
return data
raise StopIteration
class WSGIErrorsWrapper(io.RawIOBase):
def __init__(self, cfg):
# There is no public __init__ method for RawIOBase so
# we don't need to call super() in the __init__ method.
# pylint: disable=super-init-not-called
errorlog = logging.getLogger("gunicorn.error")
handlers = errorlog.handlers
self.streams = []
if cfg.errorlog == "-":
self.streams.append(sys.stderr)
handlers = handlers[1:]
for h in handlers:
if hasattr(h, "stream"):
self.streams.append(h.stream)
def write(self, data):
for stream in self.streams:
try:
stream.write(data)
except UnicodeError:
stream.write(data.encode("UTF-8"))
stream.flush()
def base_environ(cfg):
return {
"wsgi.errors": WSGIErrorsWrapper(cfg),
"wsgi.version": (1, 0),
"wsgi.multithread": False,
"wsgi.multiprocess": (cfg.workers > 1),
"wsgi.run_once": False,
"wsgi.file_wrapper": FileWrapper,
"wsgi.input_terminated": True,
"SERVER_SOFTWARE": SERVER_SOFTWARE,
}
def default_environ(req, sock, cfg):
env = base_environ(cfg)
env.update({
"wsgi.input": req.body,
"gunicorn.socket": sock,
"REQUEST_METHOD": req.method,
"QUERY_STRING": req.query,
"RAW_URI": req.uri,
"SERVER_PROTOCOL": "HTTP/%s" % ".".join([str(v) for v in req.version])
})
return env
def proxy_environ(req):
info = req.proxy_protocol_info
if not info:
return {}
return {
"PROXY_PROTOCOL": info["proxy_protocol"],
"REMOTE_ADDR": info["client_addr"],
"REMOTE_PORT": str(info["client_port"]),
"PROXY_ADDR": info["proxy_addr"],
"PROXY_PORT": str(info["proxy_port"]),
}
def _make_early_hints_callback(req, sock, resp):
"""Create a wsgi.early_hints callback for sending 103 Early Hints.
This allows WSGI applications to send 103 Early Hints responses
before the final response, enabling browsers to preload resources.
Args:
req: The request object
sock: The socket to write to
resp: The Response object to check if headers have been sent
Returns:
A callback function that accepts a list of (name, value) header tuples
and sends a 103 Early Hints response.
Note:
- Early hints are only sent for HTTP/1.1 or later clients
- HTTP/1.0 clients will silently ignore the callback
- Multiple calls are allowed (sending multiple 103 responses)
- Calls after response has started are silently ignored
"""
def send_early_hints(headers):
"""Send 103 Early Hints response.
Args:
headers: List of (name, value) header tuples, typically Link headers
Example: [('Link', '</style.css>; rel=preload; as=style')]
Raises:
InvalidHeaderName: If a header name is not a valid HTTP token.
InvalidHeader: If a header value contains invalid characters.
"""
# Don't send after response has started - would break framing
if resp.headers_sent:
return
# Don't send to HTTP/1.0 clients - they don't support 1xx responses
if req.version < (1, 1):
return
# Build 103 response
response = b"HTTP/1.1 103 Early Hints\r\n"
for name, value in headers:
if isinstance(name, bytes):
name = name.decode('latin-1')
if isinstance(value, bytes):
value = value.decode('latin-1')
# Validate header name and value using the same checks as
# Response.process_headers — defense-in-depth against
# HTTP response splitting via CRLF injection.
if not TOKEN_RE.fullmatch(name):
raise InvalidHeaderName('%r' % name)
if not HEADER_VALUE_RE.fullmatch(value):
# Pass only the name — the invalid value may contain
# sensitive data that shouldn't cross security boundaries
# via exception propagation (browsers/proxies may forward
# it to untrusted parties).
raise InvalidHeader('%r' % name)
value = value.strip(" \t")
response += f"{name}: {value}\r\n".encode('latin-1')
response += b"\r\n"
util.write(sock, response)
return send_early_hints
def create(req, sock, client, server, cfg, response_class=None,
response_args=()):
"""Build the (response, environ) pair for a request.
``response_class`` and ``response_args`` let a protocol supply its own
writer: HTTP/2 passes HTTP2Response so the body is framed as HTTP/2
instead of HTTP/1, while everything else here stays the same.
"""
if response_class is None:
response_class = Response
resp = response_class(req, sock, cfg, *response_args)
# set initial environ
environ = default_environ(req, sock, cfg)
# default variables
host = None
script_name = os.environ.get("SCRIPT_NAME", "")
if req._expected_100_continue:
sock.send(b"HTTP/1.1 100 Continue\r\n\r\n")
# rfc9112: Expect MUST be forwarded if the request is forwarded
# N.B. gunicorn just sends at most one - application might send another
# add the headers to the environ
for hdr_name, hdr_value in req.headers:
if hdr_name == 'HOST':
host = hdr_value
elif hdr_name == "SCRIPT_NAME":
script_name = hdr_value
elif hdr_name == "CONTENT-TYPE":
environ['CONTENT_TYPE'] = hdr_value
continue
elif hdr_name == "CONTENT-LENGTH":
environ['CONTENT_LENGTH'] = hdr_value
continue
# do not change lightly, this is a common source of security problems
# RFC9110 Section 17.10 discourages ambiguous or incomplete mappings
key = 'HTTP_' + hdr_name.replace('-', '_')
if key in environ:
hdr_value = "%s,%s" % (environ[key], hdr_value)
environ[key] = hdr_value
# set the url scheme
environ['wsgi.url_scheme'] = req.scheme
# set the REMOTE_* keys in environ
# authors should be aware that REMOTE_HOST and REMOTE_ADDR
# may not qualify the remote addr:
# http://www.ietf.org/rfc/rfc3875
if isinstance(client, str):
environ['REMOTE_ADDR'] = client
elif isinstance(client, bytes):
environ['REMOTE_ADDR'] = client.decode()
else:
environ['REMOTE_ADDR'] = client[0]
environ['REMOTE_PORT'] = str(client[1])
# handle the SERVER_*
# Normally only the application should use the Host header but since the
# WSGI spec doesn't support unix sockets, we are using it to create
# viable SERVER_* if possible.
if isinstance(server, str):
server = server.split(":")
if len(server) == 1:
# unix socket
if host:
server = host.split(':')
if len(server) == 1:
if req.scheme == "http":
server.append(80)
elif req.scheme == "https":
server.append(443)
else:
server.append('')
else:
# no host header given which means that we are not behind a
# proxy, so append an empty port.
server.append('')
environ['SERVER_NAME'] = server[0]
environ['SERVER_PORT'] = str(server[1])
# set the path and script name
path_info = req.path
if script_name:
if not path_info.startswith(script_name):
raise ConfigurationProblem(
"Request path %r does not start with SCRIPT_NAME %r" %
(path_info, script_name))
path_info = path_info[len(script_name):]
environ['PATH_INFO'] = util.unquote_to_wsgi_str(path_info)
environ['SCRIPT_NAME'] = script_name
# override the environ with the correct remote and server address if
# we are behind a proxy using the proxy protocol.
environ.update(proxy_environ(req))
# Add wsgi.early_hints callback for sending 103 Early Hints
environ['wsgi.early_hints'] = _make_early_hints_callback(req, sock, resp)
# Add HTTP/2 stream priority if available
if hasattr(req, 'priority_weight'):
environ['gunicorn.http2.priority_weight'] = req.priority_weight
environ['gunicorn.http2.priority_depends_on'] = req.priority_depends_on
return resp, environ
class Response:
def __init__(self, req, sock, cfg):
self.req = req
self.sock = sock
self.version = SERVER
self.status = None
self.chunked = False
self.must_close = False
self.headers = []
self.headers_sent = False
self.response_length = None
self.sent = 0
self.upgrade = False
self.cfg = cfg
self._omits_body = False
self._omits_body_warned = False
def force_close(self):
self.must_close = True
def should_close(self):
if self.must_close or self.req.should_close():
return True
if self.response_length is not None or self.chunked:
return False
if self.req.method == 'HEAD':
return False
if self.status_code < 200 or self.status_code in (204, 304):
return False
return True
def start_response(self, status, headers, exc_info=None):
if exc_info:
try:
if self.status and self.headers_sent:
util.reraise(exc_info[0], exc_info[1], exc_info[2])
finally:
exc_info = None
elif self.status is not None:
raise AssertionError("Response headers already set!")
self.status = status
# get the status code from the response here so we can use it to check
# the need for the connection header later without parsing the string
# each time.
try:
self.status_code = int(self.status.split()[0])
except ValueError:
self.status_code = None
self.process_headers(headers)
self._omits_body = self._response_omits_body(
self.req.method, self.status_code)
if self._omits_body and self._response_forbids_content_length(
self.status_code):
self.headers = [
(k, v) for k, v in self.headers if k.lower() != "content-length"
]
self.response_length = None
self.chunked = self.is_chunked()
return self.write
@staticmethod
def _response_omits_body(method, status):
# RFC 9110: HEAD requests and 1xx/204/304 responses MUST NOT carry
# a body, regardless of what the application emits.
return (
method == "HEAD"
or status in (204, 304)
or (status is not None and 100 <= status < 200)
)
@staticmethod
def _response_forbids_content_length(status):
# RFC 9110 §6.4.2: a server MUST NOT send Content-Length on 1xx or
# 204. HEAD MAY include the Content-Length the same GET would carry,
# and 304 MAY include the Content-Length of the unconditional response.
return status == 204 or (status is not None and 100 <= status < 200)
def process_headers(self, headers):
for name, value in headers:
if not isinstance(name, str):
raise TypeError('%r is not a string' % name)
if not TOKEN_RE.fullmatch(name):
raise InvalidHeaderName('%r' % name)
if not isinstance(value, str):
raise TypeError('%r is not a string' % value)
if not HEADER_VALUE_RE.fullmatch(value):
raise InvalidHeader('%r' % value)
# RFC9110 5.5
value = value.strip(" \t")
lname = name.lower()
if lname == "content-length":
self.response_length = int(value)
elif util.is_hoppish(name):
if lname == "connection":
# handle websocket
if value.lower() == "upgrade":
self.upgrade = True
elif lname == "upgrade":
if value.lower() == "websocket":
self.headers.append((name, value))
# ignore hopbyhop headers
continue
self.headers.append((name, value))
def is_chunked(self):
# Only use chunked responses when the client is
# speaking HTTP/1.1 or newer and there was
# no Content-Length header set.
if self.response_length is not None:
return False
elif self.req.version <= (1, 0):
return False
elif self._omits_body:
# No body permitted (HEAD or 1xx/204/304), so no chunked framing.
return False
return True
def default_headers(self):
# set the connection header
if self.upgrade:
connection = "upgrade"
elif self.should_close():
connection = "close"
else:
connection = "keep-alive"
headers = [
"HTTP/%s.%s %s\r\n" % (self.req.version[0],
self.req.version[1], self.status),
"Server: %s\r\n" % self.version,
"Date: %s\r\n" % util.http_date(),
"Connection: %s\r\n" % connection
]
if self.chunked:
headers.append("Transfer-Encoding: chunked\r\n")
return headers
def send_headers(self):
if self.headers_sent:
return
tosend = self.default_headers()
tosend.extend(["%s: %s\r\n" % (k, v) for k, v in self.headers])
header_str = "%s\r\n" % "".join(tosend)
util.write(self.sock, util.to_bytestring(header_str, "latin-1"))
self.headers_sent = True
def write(self, arg):
self.send_headers()
if not isinstance(arg, bytes):
raise TypeError('%r is not a byte' % arg)
if self._omits_body:
if arg and not self._omits_body_warned:
log.warning(
"WSGI app sent body bytes on a no-body response "
"(method=%s status=%s); dropping per RFC 9110.",
self.req.method, self.status_code,
)
self._omits_body_warned = True
return
arglen = len(arg)
tosend = arglen
if self.response_length is not None:
if self.sent >= self.response_length:
# Never write more than self.response_length bytes
return
tosend = min(self.response_length - self.sent, tosend)
if tosend < arglen:
arg = arg[:tosend]
# Sending an empty chunk signals the end of the
# response and prematurely closes the response
if self.chunked and tosend == 0:
return
self.sent += tosend
self._emit_body(arg)
def _emit_body(self, data):
"""Put body bytes on the wire.
The one place body framing happens, so a subclass can frame it
differently without reimplementing write()'s bookkeeping.
"""
util.write(self.sock, data, self.chunked)
def can_sendfile(self):
return self.cfg.sendfile is not False
def sendfile(self, respiter):
if self.cfg.is_ssl or not self.can_sendfile():
return False
if not util.has_fileno(respiter.filelike):
return False
fileno = respiter.filelike.fileno()
try:
offset = os.lseek(fileno, 0, os.SEEK_CUR)
if self.response_length is None:
filesize = os.fstat(fileno).st_size
nbytes = filesize - offset
else:
nbytes = self.response_length
except (OSError, io.UnsupportedOperation):
return False
if self._omits_body:
self.send_headers()
# Only complain when there really are body bytes to drop, the
# same way write() only warns for a non-empty argument.
if nbytes > 0 and not self._omits_body_warned:
log.warning(
"WSGI app sent body bytes on a no-body response "
"(method=%s status=%s); dropping per RFC 9110.",
self.req.method, self.status_code,
)
self._omits_body_warned = True
return True
self.send_headers()
if self.is_chunked():
chunk_size = "%X\r\n" % nbytes
self.sock.sendall(chunk_size.encode('utf-8'))
if nbytes > 0:
self.sock.sendfile(respiter.filelike, offset=offset, count=nbytes)
if self.is_chunked():
self.sock.sendall(b"\r\n")
os.lseek(fileno, offset, os.SEEK_SET)
return True
def write_file(self, respiter):
if not self.sendfile(respiter):
for item in respiter:
self.write(item)
def close(self):
if not self.headers_sent:
self.send_headers()
if self.chunked:
util.write_chunk(self.sock, b"")
@@ -0,0 +1,86 @@
# -*- coding: utf-8 -
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
HTTP/2 support for Gunicorn.
This module provides HTTP/2 protocol support using the hyper-h2 library.
HTTP/2 requires TLS with ALPN negotiation.
"""
H2_MIN_VERSION = (4, 1, 0)
_h2_available = None
_h2_version = None
def is_http2_available():
"""Check if HTTP/2 support is available.
Returns:
bool: True if the h2 library is installed with minimum required version.
"""
global _h2_available, _h2_version # pylint: disable=global-statement
if _h2_available is not None:
return _h2_available
try:
import h2
version_str = getattr(h2, '__version__', '0.0.0')
version_parts = tuple(int(x) for x in version_str.split('.')[:3])
_h2_version = version_parts
_h2_available = version_parts >= H2_MIN_VERSION
except ImportError:
_h2_available = False
_h2_version = None
return _h2_available
def get_h2_version():
"""Get the installed h2 library version.
Returns:
tuple: Version tuple (major, minor, patch) or None if not installed.
"""
if _h2_version is None:
is_http2_available() # Populate _h2_version
return _h2_version
def get_http2_connection_class():
"""Get the HTTP2ServerConnection class if h2 is available.
Returns:
HTTP2ServerConnection class, or raises HTTP2NotAvailable
"""
if not is_http2_available():
from .errors import HTTP2NotAvailable
raise HTTP2NotAvailable()
from .connection import HTTP2ServerConnection
return HTTP2ServerConnection
def get_async_http2_connection_class():
"""Get the AsyncHTTP2Connection class if h2 is available.
Returns:
AsyncHTTP2Connection class, or raises HTTP2NotAvailable
"""
if not is_http2_available():
from .errors import HTTP2NotAvailable
raise HTTP2NotAvailable()
from .async_connection import AsyncHTTP2Connection
return AsyncHTTP2Connection
__all__ = [
'is_http2_available',
'get_h2_version',
'get_http2_connection_class',
'get_async_http2_connection_class',
'H2_MIN_VERSION',
]
@@ -0,0 +1,656 @@
# -*- coding: utf-8 -
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
Async HTTP/2 server connection implementation for ASGI workers.
Uses the hyper-h2 library for HTTP/2 protocol handling with
asyncio for non-blocking I/O.
"""
import asyncio
import collections
from .errors import (
HTTP2Error, HTTP2ProtocolError, HTTP2ConnectionError,
HTTP2NotAvailable, HTTP2ErrorCode,
)
from .stream import HTTP2Stream
from .request import HTTP2Request
# Import h2 lazily to allow graceful fallback
_h2 = None
_h2_config = None
_h2_events = None
_h2_exceptions = None
_h2_settings = None
def _import_h2():
"""Lazily import h2 library components."""
global _h2, _h2_config, _h2_events, _h2_exceptions, _h2_settings # pylint: disable=global-statement
if _h2 is not None:
return
try:
import h2.connection as _h2
import h2.config as _h2_config
import h2.events as _h2_events
import h2.exceptions as _h2_exceptions
import h2.settings as _h2_settings
except ImportError:
raise HTTP2NotAvailable()
class AsyncHTTP2Connection:
"""Async HTTP/2 server-side connection handler for ASGI.
Manages the HTTP/2 connection state and multiplexed streams
using asyncio for non-blocking I/O operations.
"""
# Default buffer size for socket reads
READ_BUFFER_SIZE = 65536
def __init__(self, cfg, reader, writer, client_addr):
"""Initialize an async HTTP/2 server connection.
Args:
cfg: Gunicorn configuration object
reader: asyncio StreamReader
writer: asyncio StreamWriter
client_addr: Client address tuple (host, port)
Raises:
HTTP2NotAvailable: If h2 library is not installed
"""
_import_h2()
self.cfg = cfg
self.reader = reader
self.writer = writer
self.client_addr = client_addr
# Active streams indexed by stream ID
self.streams = {}
# Events pulled off the wire while blocked on a flow-control window.
# They have left the h2 state machine already, so they are held here
# for the main receive loop rather than discarded.
self._deferred_events = collections.deque()
# Queue of completed requests for the worker
self._request_queue = asyncio.Queue()
# Connection settings from config
self.initial_window_size = cfg.http2_initial_window_size
self.max_concurrent_streams = cfg.http2_max_concurrent_streams
self.max_frame_size = cfg.http2_max_frame_size
self.max_header_list_size = cfg.http2_max_header_list_size
# Initialize h2 connection
config = _h2_config.H2Configuration(
client_side=False,
header_encoding='utf-8',
)
self.h2_conn = _h2.H2Connection(config=config)
# Connection state
self._closed = False
self._initialized = False
self._receive_task = None
async def initiate_connection(self):
"""Send initial HTTP/2 settings to client.
Should be called after the SSL handshake completes and
before processing any data.
"""
if self._initialized:
return
# Update local settings before initiating
self.h2_conn.update_settings({
_h2_settings.SettingCodes.MAX_CONCURRENT_STREAMS: self.max_concurrent_streams,
_h2_settings.SettingCodes.INITIAL_WINDOW_SIZE: self.initial_window_size,
_h2_settings.SettingCodes.MAX_FRAME_SIZE: self.max_frame_size,
_h2_settings.SettingCodes.MAX_HEADER_LIST_SIZE: self.max_header_list_size,
})
self.h2_conn.initiate_connection()
await self._send_pending_data()
self._initialized = True
async def initiate_upgrade(self, settings_header, http1_req, body=b""):
"""Switch a connection to HTTP/2 after an Upgrade: h2c request.
The async twin of HTTP2Connection.initiate_upgrade. The upgraded
request becomes stream 1 (RFC 7540 section 3.2): h2 opens it in the
state machine, and the matching gunicorn stream is built here from the
HTTP/1 request that carried the upgrade, so the worker sees an
ordinary HTTP/2 request.
The body is passed in rather than read off the request: the callback
parser hands body chunks to the protocol, not to the request object.
Returns the HTTP2Request for stream 1.
"""
self.h2_conn.update_settings({
_h2_settings.SettingCodes.MAX_CONCURRENT_STREAMS: self.max_concurrent_streams,
_h2_settings.SettingCodes.INITIAL_WINDOW_SIZE: self.initial_window_size,
_h2_settings.SettingCodes.MAX_FRAME_SIZE: self.max_frame_size,
_h2_settings.SettingCodes.MAX_HEADER_LIST_SIZE: self.max_header_list_size,
})
self.h2_conn.initiate_upgrade_connection(settings_header=settings_header)
await self._send_pending_data()
self._initialized = True
stream = HTTP2Stream(stream_id=1, connection=self)
authority = ""
for name, value in http1_req.headers:
if name == "HOST":
authority = value
break
pseudo = [
(':method', http1_req.method),
(':path', http1_req.uri),
(':scheme', http1_req.scheme),
]
if authority:
pseudo.append((':authority', authority))
regular = [(name.lower(), value) for name, value in http1_req.headers
if name not in ("CONNECTION", "UPGRADE", "HTTP2-SETTINGS")]
stream.receive_headers(pseudo + regular, end_stream=not body)
if body:
stream.receive_data(body, end_stream=True)
self.streams[1] = stream
return HTTP2Request(stream, self.cfg, self.client_addr)
async def receive_data(self, timeout=None):
"""Receive data and return completed requests.
Args:
timeout: Optional timeout in seconds for read operation
Returns:
list: List of HTTP2Request objects for completed requests
Raises:
HTTP2ConnectionError: On protocol or connection errors
asyncio.TimeoutError: If timeout expires
"""
try:
if timeout is not None:
data = await asyncio.wait_for(
self.reader.read(self.READ_BUFFER_SIZE),
timeout=timeout
)
else:
data = await self.reader.read(self.READ_BUFFER_SIZE)
except (OSError, IOError) as e:
raise HTTP2ConnectionError(f"Socket read error: {e}")
if not data:
# Connection closed by peer
self._closed = True
return []
# Feed data to h2
# Note: Specific exceptions must come before ProtocolError (their parent class)
try:
events = self.h2_conn.receive_data(data)
except _h2_exceptions.FlowControlError as e:
# Send GOAWAY with FLOW_CONTROL_ERROR
await self.close(error_code=HTTP2ErrorCode.FLOW_CONTROL_ERROR)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.FrameTooLargeError as e:
# Send GOAWAY with FRAME_SIZE_ERROR
await self.close(error_code=HTTP2ErrorCode.FRAME_SIZE_ERROR)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.InvalidSettingsValueError as e:
# Use error_code from h2 exception (RFC 7540 Section 6.5.2):
# INITIAL_WINDOW_SIZE > 2^31-1 gives FLOW_CONTROL_ERROR
# Other invalid settings give PROTOCOL_ERROR
error_code = getattr(e, 'error_code', None)
if error_code is not None:
await self.close(error_code=error_code)
else:
await self.close(error_code=HTTP2ErrorCode.PROTOCOL_ERROR)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.TooManyStreamsError as e:
# Send GOAWAY with REFUSED_STREAM
await self.close(error_code=HTTP2ErrorCode.REFUSED_STREAM)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.ProtocolError as e:
# Send GOAWAY with PROTOCOL_ERROR before raising
await self.close(error_code=HTTP2ErrorCode.PROTOCOL_ERROR)
raise HTTP2ProtocolError(str(e))
# Process events, oldest first: anything set aside during a
# flow-control wait arrived before this batch.
completed_requests = []
if self._deferred_events:
events = list(self._deferred_events) + list(events)
self._deferred_events.clear()
for event in events:
request = self._handle_event(event)
if request is not None:
completed_requests.append(request)
# Send any pending data (WINDOW_UPDATE, etc.)
await self._send_pending_data()
return completed_requests
def _handle_event(self, event):
"""Handle a single h2 event.
Args:
event: h2 event object
Returns:
HTTP2Request if a request is complete, None otherwise
"""
if isinstance(event, _h2_events.RequestReceived):
return self._handle_request_received(event)
elif isinstance(event, _h2_events.DataReceived):
return self._handle_data_received(event)
elif isinstance(event, _h2_events.StreamEnded):
return self._handle_stream_ended(event)
elif isinstance(event, _h2_events.StreamReset):
self._handle_stream_reset(event)
elif isinstance(event, _h2_events.WindowUpdated):
pass # Flow control update, handled by h2
elif isinstance(event, _h2_events.PriorityUpdated):
self._handle_priority_updated(event)
elif isinstance(event, _h2_events.SettingsAcknowledged):
pass # Settings ACK received
elif isinstance(event, _h2_events.ConnectionTerminated):
self._handle_connection_terminated(event)
elif isinstance(event, _h2_events.TrailersReceived):
return self._handle_trailers_received(event)
return None
def _handle_request_received(self, event):
"""Handle RequestReceived event (HEADERS frame)."""
stream_id = event.stream_id
headers = event.headers
# Create new stream
stream = HTTP2Stream(stream_id, self)
self.streams[stream_id] = stream
# Process headers
stream.receive_headers(headers, end_stream=False)
def _handle_data_received(self, event):
"""Handle DataReceived event."""
stream_id = event.stream_id
data = event.data
stream = self.streams.get(stream_id)
if stream is None:
return None
stream.receive_data(data, end_stream=False)
# Increment flow control windows (only if data received)
if len(data) > 0:
try:
# Update stream-level window
self.h2_conn.increment_flow_control_window(len(data), stream_id=stream_id)
# Update connection-level window
self.h2_conn.increment_flow_control_window(len(data), stream_id=None)
except (ValueError, _h2_exceptions.FlowControlError):
# Window overflow - prepare GOAWAY with FLOW_CONTROL_ERROR
# (will be sent by receive_data's _send_pending_data call)
self._closed = True
try:
self.h2_conn.close_connection(error_code=HTTP2ErrorCode.FLOW_CONTROL_ERROR)
except Exception:
pass
return None
def _handle_stream_ended(self, event):
"""Handle StreamEnded event."""
stream_id = event.stream_id
stream = self.streams.get(stream_id)
if stream is None:
return None
# Mark stream as request complete and body complete so the
# receive() closure's _body_complete guard fires, preventing
# the fast path from re-reading already-consumed data from BytesIO.
stream.request_complete = True
stream._body_complete = True
if stream._body_event:
stream._body_event.set()
return HTTP2Request(stream, self.cfg, self.client_addr)
def _handle_stream_reset(self, event):
"""Handle StreamReset event."""
stream_id = event.stream_id
stream = self.streams.get(stream_id)
if stream is not None:
stream.reset(event.error_code)
def _handle_connection_terminated(self, event):
"""Handle ConnectionTerminated event."""
self._closed = True
def _handle_trailers_received(self, event):
"""Handle TrailersReceived event."""
stream_id = event.stream_id
stream = self.streams.get(stream_id)
if stream is None:
return None
stream.receive_trailers(event.headers)
return HTTP2Request(stream, self.cfg, self.client_addr)
def _handle_priority_updated(self, event):
"""Handle PriorityUpdated event (PRIORITY frame).
Args:
event: PriorityUpdated event with priority info
"""
stream = self.streams.get(event.stream_id)
if stream is not None:
stream.update_priority(
weight=event.weight,
depends_on=event.depends_on,
exclusive=event.exclusive
)
async def send_informational(self, stream_id, status, headers):
"""Send an informational response (1xx) on a stream.
This is used for 103 Early Hints and other 1xx responses.
Informational responses are sent before the final response
and do not end the stream.
Args:
stream_id: The stream ID
status: HTTP status code (100-199)
headers: List of (name, value) header tuples
Raises:
HTTP2Error: If status is not in 1xx range
"""
if status < 100 or status >= 200:
raise HTTP2Error(f"Invalid informational status: {status}")
stream = self.streams.get(stream_id)
if stream is None:
raise HTTP2Error(f"Stream {stream_id} not found")
# Build headers with :status pseudo-header
response_headers = [(':status', str(status))]
for name, value in headers:
# HTTP/2 headers must be lowercase
response_headers.append((name.lower(), str(value)))
# Send headers with end_stream=False (informational, more to follow)
self.h2_conn.send_headers(stream_id, response_headers, end_stream=False)
await self._send_pending_data()
async def send_response(self, stream_id, status, headers, body=None):
"""Send a response on a stream.
Args:
stream_id: The stream ID to respond on
status: HTTP status code (int)
headers: List of (name, value) header tuples
body: Optional response body bytes
Returns:
bool: True if response sent, False if stream was already closed
"""
stream = self.streams.get(stream_id)
if stream is None:
# Stream was already cleaned up (reset/closed) - return gracefully
return False
# Build response headers with :status pseudo-header
response_headers = [(':status', str(status))]
for name, value in headers:
response_headers.append((name.lower(), str(value)))
end_stream = body is None or len(body) == 0
try:
# Send headers
self.h2_conn.send_headers(stream_id, response_headers, end_stream=end_stream)
stream.send_headers(response_headers, end_stream=end_stream)
await self._send_pending_data()
# Send body if present
if body and len(body) > 0:
await self.send_data(stream_id, body, end_stream=True)
return True
except _h2_exceptions.StreamClosedError:
# Stream was reset by client - clean up gracefully
stream.close()
self.cleanup_stream(stream_id)
return False
async def _wait_for_flow_control_window(self, stream_id):
"""Wait for flow control window to become positive.
Returns:
int: Available window size, or -1 if waiting failed
"""
max_wait_attempts = 50 # ~5 seconds at 100ms per attempt
for _ in range(max_wait_attempts):
available = self.h2_conn.local_flow_control_window(stream_id)
if available > 0:
return available
# Read more data from connection (may receive WINDOW_UPDATE)
try:
incoming = await asyncio.wait_for(
self.reader.read(self.READ_BUFFER_SIZE),
timeout=0.1
)
if incoming:
events = self.h2_conn.receive_data(incoming)
# Process events but don't create new requests
for event in events:
if isinstance(event, _h2_events.StreamReset):
if event.stream_id == stream_id:
return -1
elif isinstance(event, _h2_events.ConnectionTerminated):
self._closed = True
return -1
else:
# Anything else arriving alongside the
# WINDOW_UPDATE belongs to the main loop. It has
# already left the h2 state machine, so dropping
# it here loses a request or its body for good.
self._deferred_events.append(event)
await self._send_pending_data()
else:
# Connection closed
self._closed = True
return -1
except asyncio.TimeoutError:
continue
except _h2_exceptions.ProtocolError:
return -1
return self.h2_conn.local_flow_control_window(stream_id)
async def send_data(self, stream_id, data, end_stream=False):
"""Send data on a stream.
Args:
stream_id: The stream ID
data: Body data bytes
end_stream: Whether this ends the stream
Returns:
bool: True if data sent, False if stream was already closed
"""
stream = self.streams.get(stream_id)
if stream is None:
return False
data_to_send = data
try:
while data_to_send:
available = self.h2_conn.local_flow_control_window(stream_id)
chunk_size = min(available, self.max_frame_size, len(data_to_send))
if chunk_size <= 0:
# Wait for WINDOW_UPDATE per RFC 7540 Section 6.9.2
await self._send_pending_data()
available = await self._wait_for_flow_control_window(stream_id)
if available <= 0:
return False
chunk_size = min(available, self.max_frame_size, len(data_to_send))
chunk = data_to_send[:chunk_size]
data_to_send = data_to_send[chunk_size:]
is_final = end_stream and len(data_to_send) == 0
self.h2_conn.send_data(stream_id, chunk, end_stream=is_final)
await self._send_pending_data()
stream.send_data(data, end_stream=end_stream)
return True
except (_h2_exceptions.StreamClosedError, _h2_exceptions.FlowControlError):
stream.close()
self.cleanup_stream(stream_id)
return False
async def send_trailers(self, stream_id, trailers):
"""Send trailing headers on a stream.
Trailers are headers sent after the response body, commonly used
for gRPC status codes, checksums, and timing information.
Args:
stream_id: The stream ID
trailers: List of (name, value) trailer tuples
Raises:
HTTP2Error: If stream not found, headers not sent, or pseudo-headers used
Returns:
bool: True if trailers sent, False if stream was already closed
"""
stream = self.streams.get(stream_id)
if stream is None:
# Stream was already cleaned up (reset/closed) - return gracefully
return False
if not stream.response_headers_sent:
# Can't send trailers without headers - return False
return False
# Validate and normalize trailer headers
trailer_headers = []
for name, value in trailers:
lname = name.lower()
if lname.startswith(':'):
raise HTTP2Error(f"Pseudo-header '{name}' not allowed in trailers")
trailer_headers.append((lname, str(value)))
try:
# Send trailers with end_stream=True
self.h2_conn.send_headers(stream_id, trailer_headers, end_stream=True)
stream.send_trailers(trailer_headers)
await self._send_pending_data()
return True
except _h2_exceptions.StreamClosedError:
# Stream was reset by client - clean up gracefully
stream.close()
self.cleanup_stream(stream_id)
return False
async def send_error(self, stream_id, status_code, message=None):
"""Send an error response on a stream."""
body = message.encode() if message else b''
headers = [('content-length', str(len(body)))]
if body:
headers.append(('content-type', 'text/plain; charset=utf-8'))
await self.send_response(stream_id, status_code, headers, body)
async def reset_stream(self, stream_id, error_code=0x8):
"""Reset a stream with RST_STREAM."""
stream = self.streams.get(stream_id)
if stream is not None:
stream.reset(error_code)
self.h2_conn.reset_stream(stream_id, error_code=error_code)
await self._send_pending_data()
async def close(self, error_code=0x0, last_stream_id=None):
"""Close the connection gracefully with GOAWAY."""
if self._closed:
return
self._closed = True
if last_stream_id is None:
last_stream_id = max(self.streams.keys()) if self.streams else 0
try:
self.h2_conn.close_connection(error_code=error_code)
await self._send_pending_data()
except Exception:
pass
try:
self.writer.close()
await self.writer.wait_closed()
except Exception:
pass
async def _send_pending_data(self):
"""Send any pending data from h2 to the socket."""
data = self.h2_conn.data_to_send()
if data:
try:
self.writer.write(data)
await self.writer.drain()
except (OSError, IOError) as e:
self._closed = True
raise HTTP2ConnectionError(f"Socket write error: {e}")
@property
def is_closed(self):
"""Check if connection is closed."""
return self._closed
def cleanup_stream(self, stream_id):
"""Remove a stream after processing is complete."""
self.streams.pop(stream_id, None)
def __repr__(self):
return (
f"<AsyncHTTP2Connection "
f"streams={len(self.streams)} "
f"closed={self._closed}>"
)
__all__ = ['AsyncHTTP2Connection']
@@ -0,0 +1,765 @@
# -*- coding: utf-8 -
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
HTTP/2 server connection implementation.
Uses the hyper-h2 library for HTTP/2 protocol handling.
"""
import collections
import selectors
from io import BytesIO
from .errors import (
HTTP2Error, HTTP2ProtocolError, HTTP2ConnectionError,
HTTP2NotAvailable, HTTP2ErrorCode,
)
from .stream import HTTP2Stream
from .request import HTTP2Request
# Import h2 lazily to allow graceful fallback
_h2 = None
_h2_config = None
_h2_events = None
_h2_exceptions = None
_h2_settings = None
def _import_h2():
"""Lazily import h2 library components."""
global _h2, _h2_config, _h2_events, _h2_exceptions, _h2_settings # pylint: disable=global-statement
if _h2 is not None:
return
try:
import h2.connection as _h2
import h2.config as _h2_config
import h2.events as _h2_events
import h2.exceptions as _h2_exceptions
import h2.settings as _h2_settings
except ImportError:
raise HTTP2NotAvailable()
class HTTP2ServerConnection:
"""HTTP/2 server-side connection handler.
Manages the HTTP/2 connection state and multiplexed streams.
This class wraps the h2 library and provides a higher-level
interface for gunicorn workers.
"""
# Default buffer size for socket reads
READ_BUFFER_SIZE = 65536
def __init__(self, cfg, sock, client_addr):
"""Initialize an HTTP/2 server connection.
Args:
cfg: Gunicorn configuration object
sock: SSL socket with completed handshake
client_addr: Client address tuple (host, port)
Raises:
HTTP2NotAvailable: If h2 library is not installed
"""
_import_h2()
self.cfg = cfg
self.sock = sock
self.client_addr = client_addr
# Active streams indexed by stream ID
self.streams = {}
# Events pulled off the wire while blocked on a flow-control window.
# They have left the h2 state machine already, so they are held here
# for the main receive loop rather than discarded.
self._deferred_events = collections.deque()
# Completed requests ready for processing
self._pending_requests = []
# Connection settings from config
self.initial_window_size = cfg.http2_initial_window_size
self.max_concurrent_streams = cfg.http2_max_concurrent_streams
self.max_frame_size = cfg.http2_max_frame_size
self.max_header_list_size = cfg.http2_max_header_list_size
# Initialize h2 connection
config = _h2_config.H2Configuration(
client_side=False,
header_encoding='utf-8',
)
self.h2_conn = _h2.H2Connection(config=config)
# Read buffer for partial frames
self._read_buffer = BytesIO()
# Connection state
self._closed = False
self._initialized = False
def initiate_connection(self):
"""Send initial HTTP/2 settings to client.
Should be called after the SSL handshake completes and
before processing any data.
"""
if self._initialized:
return
# Update local settings before initiating
self.h2_conn.update_settings({
_h2_settings.SettingCodes.MAX_CONCURRENT_STREAMS: self.max_concurrent_streams,
_h2_settings.SettingCodes.INITIAL_WINDOW_SIZE: self.initial_window_size,
_h2_settings.SettingCodes.MAX_FRAME_SIZE: self.max_frame_size,
_h2_settings.SettingCodes.MAX_HEADER_LIST_SIZE: self.max_header_list_size,
})
self.h2_conn.initiate_connection()
self._send_pending_data()
self._initialized = True
def initiate_upgrade(self, settings_header, http1_req, body=None):
"""Switch a connection to HTTP/2 after an Upgrade: h2c request.
The upgraded request becomes stream 1 (RFC 7540 section 3.2). h2
opens it in the state machine; the matching gunicorn stream is built
here from the HTTP/1 request that carried the upgrade, so the worker
sees an ordinary HTTP/2 request.
``body`` is the request payload. A caller that has to drain it before
collecting the bytes pipelined behind the request passes it here;
leaving it None reads it off the request.
Returns the HTTP2Request for stream 1.
"""
self.h2_conn.update_settings({
_h2_settings.SettingCodes.MAX_CONCURRENT_STREAMS: self.max_concurrent_streams,
_h2_settings.SettingCodes.INITIAL_WINDOW_SIZE: self.initial_window_size,
_h2_settings.SettingCodes.MAX_FRAME_SIZE: self.max_frame_size,
_h2_settings.SettingCodes.MAX_HEADER_LIST_SIZE: self.max_header_list_size,
})
self.h2_conn.initiate_upgrade_connection(settings_header=settings_header)
self._send_pending_data()
self._initialized = True
stream = HTTP2Stream(stream_id=1, connection=self)
authority = ""
for name, value in http1_req.headers:
if name == "HOST":
authority = value
break
pseudo = [
(':method', http1_req.method),
(':path', http1_req.uri),
(':scheme', http1_req.scheme),
]
if authority:
pseudo.append((':authority', authority))
regular = [(name.lower(), value) for name, value in http1_req.headers
if name not in ("CONNECTION", "UPGRADE", "HTTP2-SETTINGS")]
if body is None:
body = b""
if http1_req.body is not None:
body = http1_req.body.read() or b""
stream.receive_headers(pseudo + regular, end_stream=not body)
if body:
stream.receive_data(body, end_stream=True)
self.streams[1] = stream
return HTTP2Request(stream, self.cfg, self.client_addr)
def receive_data(self, data=None):
"""Process received data and return completed requests.
Args:
data: Optional bytes to process. If None, reads from socket.
Returns:
list: List of HTTP2Request objects for completed requests
Raises:
HTTP2ConnectionError: On protocol or connection errors
"""
if data is None:
try:
data = self.sock.recv(self.READ_BUFFER_SIZE)
except (OSError, IOError) as e:
raise HTTP2ConnectionError(f"Socket read error: {e}")
if not data:
# Connection closed by peer
self._closed = True
return []
# Feed data to h2
# Note: Specific exceptions must come before ProtocolError (their parent class)
try:
events = self.h2_conn.receive_data(data)
except _h2_exceptions.FlowControlError as e:
# Send GOAWAY with FLOW_CONTROL_ERROR
self.close(error_code=HTTP2ErrorCode.FLOW_CONTROL_ERROR)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.FrameTooLargeError as e:
# Send GOAWAY with FRAME_SIZE_ERROR
self.close(error_code=HTTP2ErrorCode.FRAME_SIZE_ERROR)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.InvalidSettingsValueError as e:
# Use error_code from h2 exception (RFC 7540 Section 6.5.2):
# INITIAL_WINDOW_SIZE > 2^31-1 gives FLOW_CONTROL_ERROR
# Other invalid settings give PROTOCOL_ERROR
error_code = getattr(e, 'error_code', None)
if error_code is not None:
self.close(error_code=error_code)
else:
self.close(error_code=HTTP2ErrorCode.PROTOCOL_ERROR)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.TooManyStreamsError as e:
# Send GOAWAY with REFUSED_STREAM
self.close(error_code=HTTP2ErrorCode.REFUSED_STREAM)
raise HTTP2ProtocolError(str(e))
except _h2_exceptions.ProtocolError as e:
# Send GOAWAY with PROTOCOL_ERROR before raising
self.close(error_code=HTTP2ErrorCode.PROTOCOL_ERROR)
raise HTTP2ProtocolError(str(e))
# Process events, oldest first: anything set aside during a
# flow-control wait arrived before this batch.
completed_requests = []
if self._deferred_events:
events = list(self._deferred_events) + list(events)
self._deferred_events.clear()
for event in events:
request = self._handle_event(event)
if request is not None:
completed_requests.append(request)
# Send any pending data (WINDOW_UPDATE, etc.)
self._send_pending_data()
return completed_requests
def _handle_event(self, event):
"""Handle a single h2 event.
Args:
event: h2 event object
Returns:
HTTP2Request if a request is complete, None otherwise
"""
if isinstance(event, _h2_events.RequestReceived):
return self._handle_request_received(event)
elif isinstance(event, _h2_events.DataReceived):
return self._handle_data_received(event)
elif isinstance(event, _h2_events.StreamEnded):
return self._handle_stream_ended(event)
elif isinstance(event, _h2_events.StreamReset):
self._handle_stream_reset(event)
elif isinstance(event, _h2_events.WindowUpdated):
pass # Flow control update, handled by h2
elif isinstance(event, _h2_events.PriorityUpdated):
self._handle_priority_updated(event)
elif isinstance(event, _h2_events.SettingsAcknowledged):
pass # Settings ACK received
elif isinstance(event, _h2_events.ConnectionTerminated):
self._handle_connection_terminated(event)
elif isinstance(event, _h2_events.TrailersReceived):
return self._handle_trailers_received(event)
return None
def _handle_request_received(self, event):
"""Handle RequestReceived event (HEADERS frame).
Args:
event: RequestReceived event with headers
"""
stream_id = event.stream_id
headers = event.headers
# Create new stream
stream = HTTP2Stream(stream_id, self)
self.streams[stream_id] = stream
# Process headers
# The StreamEnded event will come separately for GET/HEAD with no body
stream.receive_headers(headers, end_stream=False)
def _handle_data_received(self, event):
"""Handle DataReceived event.
Args:
event: DataReceived event with body data
Returns:
None (request completion handled by StreamEnded)
"""
stream_id = event.stream_id
data = event.data
stream = self.streams.get(stream_id)
if stream is None:
# Stream was reset or doesn't exist
return None
stream.receive_data(data, end_stream=False)
# Increment flow control windows (only if data received)
if len(data) > 0:
try:
# Update stream-level window
self.h2_conn.increment_flow_control_window(len(data), stream_id=stream_id)
# Update connection-level window
self.h2_conn.increment_flow_control_window(len(data), stream_id=None)
# Send WINDOW_UPDATE frames immediately
self._send_pending_data()
except (ValueError, _h2_exceptions.FlowControlError):
# Window overflow - send FLOW_CONTROL_ERROR and close
self.close(error_code=HTTP2ErrorCode.FLOW_CONTROL_ERROR)
return None
def _handle_stream_ended(self, event):
"""Handle StreamEnded event.
Args:
event: StreamEnded event
Returns:
HTTP2Request for the completed request
"""
stream_id = event.stream_id
stream = self.streams.get(stream_id)
if stream is None:
return None
# Mark stream as request complete and body complete so the
# receive() closure's _body_complete guard fires, preventing
# the fast path from re-reading already-consumed data from BytesIO.
stream.request_complete = True
stream._body_complete = True
if stream._body_event:
stream._body_event.set()
# Create request object
return HTTP2Request(stream, self.cfg, self.client_addr)
def _handle_stream_reset(self, event):
"""Handle StreamReset event (RST_STREAM frame).
Args:
event: StreamReset event
"""
stream_id = event.stream_id
stream = self.streams.get(stream_id)
if stream is not None:
stream.reset(event.error_code)
# Keep stream in dict for potential cleanup
def _handle_connection_terminated(self, event):
"""Handle ConnectionTerminated event (GOAWAY frame).
Args:
event: ConnectionTerminated event
"""
self._closed = True
# Could log event.error_code and event.additional_data
def _handle_trailers_received(self, event):
"""Handle TrailersReceived event.
Args:
event: TrailersReceived event with trailer headers
Returns:
HTTP2Request if this completes the request
"""
stream_id = event.stream_id
stream = self.streams.get(stream_id)
if stream is None:
return None
stream.receive_trailers(event.headers)
# Trailers always end the request
return HTTP2Request(stream, self.cfg, self.client_addr)
def _handle_priority_updated(self, event):
"""Handle PriorityUpdated event (PRIORITY frame).
Args:
event: PriorityUpdated event with priority info
"""
stream = self.streams.get(event.stream_id)
if stream is not None:
stream.update_priority(
weight=event.weight,
depends_on=event.depends_on,
exclusive=event.exclusive
)
def send_informational(self, stream_id, status, headers):
"""Send an informational response (1xx) on a stream.
This is used for 103 Early Hints and other 1xx responses.
Informational responses are sent before the final response
and do not end the stream.
Args:
stream_id: The stream ID
status: HTTP status code (100-199)
headers: List of (name, value) header tuples
Raises:
HTTP2Error: If status is not in 1xx range
"""
if status < 100 or status >= 200:
raise HTTP2Error(f"Invalid informational status: {status}")
stream = self.streams.get(stream_id)
if stream is None:
raise HTTP2Error(f"Stream {stream_id} not found")
# Build headers with :status pseudo-header
response_headers = [(':status', str(status))]
for name, value in headers:
# HTTP/2 headers must be lowercase
response_headers.append((name.lower(), str(value)))
# Send headers with end_stream=False (informational, more to follow)
self.h2_conn.send_headers(stream_id, response_headers, end_stream=False)
self._send_pending_data()
def send_response_headers(self, stream_id, status, headers,
end_stream=False):
"""Send response headers on a stream without ending it.
Returns False if the stream is already gone. Split out of
send_response() so a response can be streamed: headers first, then
any number of data frames, then end_stream().
"""
stream = self.streams.get(stream_id)
if stream is None:
# Stream was already cleaned up (reset/closed)
return False
# Build response headers with :status pseudo-header
response_headers = [(':status', str(status))]
for name, value in headers:
# HTTP/2 headers must be lowercase
response_headers.append((name.lower(), str(value)))
self.h2_conn.send_headers(stream_id, response_headers,
end_stream=end_stream)
stream.send_headers(response_headers, end_stream=end_stream)
self._send_pending_data()
return True
def end_stream(self, stream_id, trailers=None):
"""Close the sending half of a stream, with trailers if given."""
if self.streams.get(stream_id) is None:
return False
if trailers:
self.send_trailers(stream_id, trailers)
return True
# Not send_data(): it chunks against the flow-control window and an
# empty payload skips that loop entirely, so END_STREAM would never
# reach the peer and the client would wait for a response that is
# already finished.
# Not send_data(): it chunks against the flow-control window and an
# empty payload skips that loop entirely, so END_STREAM would never
# reach the peer and the client would wait for a response that is
# already finished.
try:
self.h2_conn.send_data(stream_id, b"", end_stream=True)
self.streams[stream_id].send_data(b"", end_stream=True)
self._send_pending_data()
except _h2_exceptions.StreamClosedError:
self.cleanup_stream(stream_id)
return False
return True
def send_response(self, stream_id, status, headers, body=None):
"""Send a response on a stream.
Args:
stream_id: The stream ID to respond on
status: HTTP status code (int)
headers: List of (name, value) header tuples
body: Optional response body bytes
Raises:
HTTP2Error: If stream not found or in invalid state
Returns:
bool: True if response sent, False if stream was already closed
"""
end_stream = body is None or len(body) == 0
try:
if not self.send_response_headers(stream_id, status, headers,
end_stream=end_stream):
return False
# Send body if present
if body and len(body) > 0:
self.send_data(stream_id, body, end_stream=True)
return True
except _h2_exceptions.StreamClosedError:
# Stream was reset by client - clean up gracefully
stream = self.streams.get(stream_id)
if stream is not None:
stream.close()
self.cleanup_stream(stream_id)
return False
def _wait_for_flow_control_window(self, stream_id):
"""Wait for flow control window to become positive.
Returns:
int: Available window size, or -1 if waiting failed
"""
max_wait_attempts = 50 # ~5 seconds at 100ms per attempt
try:
sel = selectors.DefaultSelector()
sel.register(self.sock, selectors.EVENT_READ)
except (TypeError, ValueError):
# Socket doesn't support selectors (e.g., mock socket)
return -1
result = -1
try:
for _ in range(max_wait_attempts):
available = self.h2_conn.local_flow_control_window(stream_id)
if available > 0:
result = available
break
ready = sel.select(timeout=0.1)
if ready:
try:
incoming = self.sock.recv(self.READ_BUFFER_SIZE)
except (OSError, IOError, _h2_exceptions.ProtocolError):
break
if not incoming:
self._closed = True
break
try:
events = self.h2_conn.receive_data(incoming)
except _h2_exceptions.ProtocolError:
break
for event in events:
if isinstance(event, _h2_events.StreamReset):
if event.stream_id == stream_id:
result = -1
break
elif isinstance(event, _h2_events.ConnectionTerminated):
self._closed = True
result = -1
break
else:
# Anything else arriving alongside the
# WINDOW_UPDATE belongs to the main loop. It has
# already left the h2 state machine, so dropping
# it here loses a request or its body for good.
self._deferred_events.append(event)
else:
self._send_pending_data()
continue
break # Break outer loop if inner loop broke
else:
# Loop completed without break - check final window
result = self.h2_conn.local_flow_control_window(stream_id)
finally:
sel.close()
return result
def send_data(self, stream_id, data, end_stream=False):
"""Send data on a stream.
Args:
stream_id: The stream ID
data: Body data bytes
end_stream: Whether this ends the stream
Returns:
bool: True if data sent, False if stream was already closed
"""
stream = self.streams.get(stream_id)
if stream is None:
return False
data_to_send = data
try:
while data_to_send:
available = self.h2_conn.local_flow_control_window(stream_id)
chunk_size = min(available, self.max_frame_size, len(data_to_send))
if chunk_size <= 0:
# Wait for WINDOW_UPDATE per RFC 7540 Section 6.9.2
self._send_pending_data()
available = self._wait_for_flow_control_window(stream_id)
if available <= 0:
return False
chunk_size = min(available, self.max_frame_size, len(data_to_send))
chunk = data_to_send[:chunk_size]
data_to_send = data_to_send[chunk_size:]
is_final = end_stream and len(data_to_send) == 0
self.h2_conn.send_data(stream_id, chunk, end_stream=is_final)
self._send_pending_data()
stream.send_data(data, end_stream=end_stream)
return True
except (_h2_exceptions.StreamClosedError, _h2_exceptions.FlowControlError):
# Stream was reset by client or flow control error - clean up gracefully
stream.close()
self.cleanup_stream(stream_id)
return False
def send_trailers(self, stream_id, trailers):
"""Send trailing headers on a stream.
Trailers are headers sent after the response body, commonly used
for gRPC status codes, checksums, and timing information.
Args:
stream_id: The stream ID
trailers: List of (name, value) trailer tuples
Raises:
HTTP2Error: If stream not found, headers not sent, or pseudo-headers used
Returns:
bool: True if trailers sent, False if stream was already closed
"""
stream = self.streams.get(stream_id)
if stream is None:
# Stream was already cleaned up (reset/closed) - return gracefully
return False
if not stream.response_headers_sent:
# Can't send trailers without headers - return False
return False
# Validate and normalize trailer headers
trailer_headers = []
for name, value in trailers:
lname = name.lower()
if lname.startswith(':'):
raise HTTP2Error(f"Pseudo-header '{name}' not allowed in trailers")
trailer_headers.append((lname, str(value)))
try:
# Send trailers with end_stream=True
self.h2_conn.send_headers(stream_id, trailer_headers, end_stream=True)
stream.send_trailers(trailer_headers)
self._send_pending_data()
return True
except _h2_exceptions.StreamClosedError:
# Stream was reset by client - clean up gracefully
stream.close()
self.cleanup_stream(stream_id)
return False
def send_error(self, stream_id, status_code, message=None):
"""Send an error response on a stream.
Args:
stream_id: The stream ID
status_code: HTTP status code
message: Optional error message body
"""
body = message.encode() if message else b''
headers = [('content-length', str(len(body)))]
if body:
headers.append(('content-type', 'text/plain; charset=utf-8'))
self.send_response(stream_id, status_code, headers, body)
def reset_stream(self, stream_id, error_code=0x8):
"""Reset a stream with RST_STREAM.
Args:
stream_id: The stream ID to reset
error_code: HTTP/2 error code (default: CANCEL)
"""
stream = self.streams.get(stream_id)
if stream is not None:
stream.reset(error_code)
self.h2_conn.reset_stream(stream_id, error_code=error_code)
self._send_pending_data()
def close(self, error_code=0x0, last_stream_id=None):
"""Close the connection gracefully with GOAWAY.
Args:
error_code: HTTP/2 error code (default: NO_ERROR)
last_stream_id: Last processed stream ID (default: highest)
"""
if self._closed:
return
self._closed = True
if last_stream_id is None:
# Use highest stream ID we've seen
last_stream_id = max(self.streams.keys()) if self.streams else 0
try:
self.h2_conn.close_connection(error_code=error_code)
self._send_pending_data()
except Exception:
pass # Best effort
def _send_pending_data(self):
"""Send any pending data from h2 to the socket."""
data = self.h2_conn.data_to_send()
if data:
try:
self.sock.sendall(data)
except (OSError, IOError) as e:
self._closed = True
raise HTTP2ConnectionError(f"Socket write error: {e}")
@property
def is_closed(self):
"""Check if connection is closed."""
return self._closed
def cleanup_stream(self, stream_id):
"""Remove a stream after processing is complete.
Args:
stream_id: The stream ID to clean up
"""
self.streams.pop(stream_id, None)
def __repr__(self):
return (
f"<HTTP2ServerConnection "
f"streams={len(self.streams)} "
f"closed={self._closed}>"
)
__all__ = ['HTTP2ServerConnection']
+169
View File
@@ -0,0 +1,169 @@
# -*- coding: utf-8 -
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
HTTP/2 specific exceptions.
These exceptions map to HTTP/2 error codes defined in RFC 7540.
"""
class HTTP2ErrorCode:
"""HTTP/2 Error Codes (RFC 7540 Section 7)."""
NO_ERROR = 0x0
PROTOCOL_ERROR = 0x1
INTERNAL_ERROR = 0x2
FLOW_CONTROL_ERROR = 0x3
SETTINGS_TIMEOUT = 0x4
STREAM_CLOSED = 0x5
FRAME_SIZE_ERROR = 0x6
REFUSED_STREAM = 0x7
CANCEL = 0x8
COMPRESSION_ERROR = 0x9
CONNECT_ERROR = 0xa
ENHANCE_YOUR_CALM = 0xb
INADEQUATE_SECURITY = 0xc
HTTP_1_1_REQUIRED = 0xd
class HTTP2Error(Exception):
"""Base exception for HTTP/2 errors."""
error_code = 0x0 # NO_ERROR
def __init__(self, message=None, error_code=None):
self.message = message or self.__class__.__doc__
if error_code is not None:
self.error_code = error_code
super().__init__(self.message)
class HTTP2ProtocolError(HTTP2Error):
"""Protocol error detected."""
error_code = 0x1 # PROTOCOL_ERROR
class HTTP2InternalError(HTTP2Error):
"""Internal error occurred."""
error_code = 0x2 # INTERNAL_ERROR
class HTTP2FlowControlError(HTTP2Error):
"""Flow control limits exceeded."""
error_code = 0x3 # FLOW_CONTROL_ERROR
class HTTP2SettingsTimeout(HTTP2Error):
"""Settings acknowledgment timeout."""
error_code = 0x4 # SETTINGS_TIMEOUT
class HTTP2StreamClosed(HTTP2Error):
"""Stream was closed."""
error_code = 0x5 # STREAM_CLOSED
class HTTP2FrameSizeError(HTTP2Error):
"""Frame size is incorrect."""
error_code = 0x6 # FRAME_SIZE_ERROR
class HTTP2RefusedStream(HTTP2Error):
"""Stream was refused."""
error_code = 0x7 # REFUSED_STREAM
class HTTP2Cancel(HTTP2Error):
"""Stream was cancelled."""
error_code = 0x8 # CANCEL
class HTTP2CompressionError(HTTP2Error):
"""Compression state error."""
error_code = 0x9 # COMPRESSION_ERROR
class HTTP2ConnectError(HTTP2Error):
"""Connection error during CONNECT."""
error_code = 0xa # CONNECT_ERROR
class HTTP2EnhanceYourCalm(HTTP2Error):
"""Peer is generating excessive load."""
error_code = 0xb # ENHANCE_YOUR_CALM
class HTTP2InadequateSecurity(HTTP2Error):
"""Transport security is inadequate."""
error_code = 0xc # INADEQUATE_SECURITY
class HTTP2RequiresHTTP11(HTTP2Error):
"""HTTP/1.1 is required for this request."""
error_code = 0xd # HTTP_1_1_REQUIRED
class HTTP2StreamError(HTTP2Error):
"""Error specific to a single stream."""
def __init__(self, stream_id, message=None, error_code=None):
self.stream_id = stream_id
super().__init__(message, error_code)
def __str__(self):
return f"Stream {self.stream_id}: {self.message}"
class HTTP2ConnectionError(HTTP2Error):
"""Error affecting the entire connection."""
class HTTP2ConfigurationError(HTTP2Error):
"""Invalid HTTP/2 configuration."""
class HTTP2NotAvailable(HTTP2Error):
"""HTTP/2 support is not available (h2 library not installed)."""
def __init__(self, message=None):
message = message or "HTTP/2 requires the h2 library: pip install gunicorn[http2]"
super().__init__(message)
__all__ = [
'HTTP2ErrorCode',
'HTTP2Error',
'HTTP2ProtocolError',
'HTTP2InternalError',
'HTTP2FlowControlError',
'HTTP2SettingsTimeout',
'HTTP2StreamClosed',
'HTTP2FrameSizeError',
'HTTP2RefusedStream',
'HTTP2Cancel',
'HTTP2CompressionError',
'HTTP2ConnectError',
'HTTP2EnhanceYourCalm',
'HTTP2InadequateSecurity',
'HTTP2RequiresHTTP11',
'HTTP2StreamError',
'HTTP2ConnectionError',
'HTTP2ConfigurationError',
'HTTP2NotAvailable',
]
@@ -0,0 +1,170 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""Cleartext HTTP/2 (h2c) negotiation, shared by every worker.
The I/O differs per worker: gthread and gevent read from a socket, the ASGI
worker is handed bytes by asyncio. The decisions do not. Keeping the pure,
I/O-free part here stops the blocking and push-based paths from drifting apart.
"""
import time
from gunicorn.http.message import _ip_in_allow_list
#: HTTP/2 connection preface sent by clients using prior knowledge,
#: RFC 9113 section 3.4.
H2C_PREFACE = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n"
#: How long to wait for the whole preface once the first bytes arrive. This
#: is a budget for the entire preface, not per read.
H2C_PREFACE_TIMEOUT = 1.0
MATCH = "match"
PARTIAL = "partial"
MISMATCH = "mismatch"
def preface_match(buf):
"""Compare buffered bytes against the connection preface.
Returns ``MATCH`` when the whole preface is present, ``PARTIAL`` when the
bytes so far are a prefix of it and more could still arrive, and
``MISMATCH`` as soon as a byte diverges. Never blocks and never reads.
"""
if len(buf) >= len(H2C_PREFACE):
return MATCH if buf.startswith(H2C_PREFACE) else MISMATCH
return PARTIAL if H2C_PREFACE.startswith(buf) else MISMATCH
def peer_trusted_for_h2c(cfg, peer_addr):
"""Whether this peer may negotiate cleartext HTTP/2.
Reuses the ``forwarded_allow_ips`` trust list: h2c is only ever expected
from the TLS-terminating proxy in front of gunicorn, which is the same
peer already trusted to set forwarded headers. Unix socket peers are
trusted, matching that policy.
"""
if not isinstance(peer_addr, tuple):
return True
return _ip_in_allow_list(
peer_addr[0], cfg.forwarded_allow_ips, cfg.forwarded_allow_networks()
)
def _h2c_available(cfg):
"""Whether cleartext HTTP/2 could apply to this server at all."""
return (
"h2" in cfg.http_protocols
and getattr(cfg, "protocol", "http") == "http"
and not cfg.is_ssl
)
def prior_knowledge_allowed(cfg, peer_addr):
"""Whether to sniff for the connection preface from this peer.
Deliberately separate from :func:`upgrade_allowed`: enabling one mechanism
must not quietly enable the other.
"""
if cfg.http2_cleartext not in ("prior-knowledge", "both"):
return False
return _h2c_available(cfg) and peer_trusted_for_h2c(cfg, peer_addr)
def mismatch_is_error(cfg):
"""Whether a trusted peer failing to send the preface is a 400.
Only when prior knowledge is the sole mechanism: such a peer is expected
to speak HTTP/2 and a silent downgrade would hide a misconfiguration.
When upgrade is also enabled, an HTTP/1 request is not a mistake, it is
how an upgrade begins, so it has to be allowed through.
"""
return cfg.http2_cleartext == "prior-knowledge"
def upgrade_allowed(cfg, peer_addr):
"""Whether to honour an ``Upgrade: h2c`` request from this peer."""
if cfg.http2_cleartext not in ("upgrade", "both"):
return False
return _h2c_available(cfg) and peer_trusted_for_h2c(cfg, peer_addr)
def read_preface_blocking(sock, timeout=None):
"""Read up to the length of the preface from a blocking socket.
Returns ``(matched, consumed_bytes)``. The caller owns the consumed bytes
and must hand them to whichever protocol wins, since they have already
left the socket.
The timeout is an absolute budget for the whole preface, checked before
every read. ``socket.settimeout()`` alone would bound each call instead,
which lets a client trickle one byte per interval and hold the connection
(and, on gthread, a pool slot) for as many intervals as the preface has
bytes.
"""
if timeout is None:
# read at call time so the module attribute stays adjustable
timeout = H2C_PREFACE_TIMEOUT
buf = b""
deadline = time.monotonic() + timeout
original = sock.gettimeout()
try:
while True:
remaining = deadline - time.monotonic()
if remaining <= 0:
return False, buf
sock.settimeout(remaining)
try:
chunk = sock.recv(len(H2C_PREFACE) - len(buf))
except (TimeoutError, OSError):
return False, buf
if not chunk:
return False, buf
buf += chunk
state = preface_match(buf)
if state is MATCH:
return True, buf
if state is MISMATCH:
return False, buf
finally:
sock.settimeout(original)
#: Sent before switching an HTTP/1.1 connection over to HTTP/2.
UPGRADE_101 = (
b"HTTP/1.1 101 Switching Protocols\r\n"
b"Connection: Upgrade\r\n"
b"Upgrade: h2c\r\n"
b"\r\n"
)
def upgrade_settings(req):
"""Return the HTTP2-Settings payload if this request asks for h2c.
RFC 7540 section 3.2: the request must name ``h2c`` in Upgrade and carry
exactly one HTTP2-Settings header, itself named in Connection. Returns
None when the request is not a well-formed upgrade attempt, so the caller
simply carries on with HTTP/1.
"""
upgrade = None
settings = []
connection = ""
for name, value in req.headers:
if name == "UPGRADE":
upgrade = value.strip().lower()
elif name == "HTTP2-SETTINGS":
settings.append(value.strip())
elif name == "CONNECTION":
connection = value.lower()
if upgrade != "h2c":
return None
# Exactly one, per RFC 7540 3.2.1: a second one is ambiguous.
if len(settings) != 1:
return None
if "http2-settings" not in connection or "upgrade" not in connection:
return None
return settings[0].encode("latin-1")
+267
View File
@@ -0,0 +1,267 @@
# -*- coding: utf-8 -
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""
HTTP/2 request wrapper.
Provides a Request-compatible interface for HTTP/2 streams.
"""
from io import BytesIO
from gunicorn.http.message import (
HeaderPolicy,
RFC9110_5_5_INVALID_AND_DANGEROUS,
)
from gunicorn.http.errors import InvalidHeader
from gunicorn.util import split_request_uri
class HTTP2Body:
"""Body wrapper for HTTP/2 request data.
Provides a file-like interface to the request body,
compatible with gunicorn's Body class expectations.
"""
def __init__(self, data):
"""Initialize with body data.
Args:
data: bytes containing the request body
"""
self._data = BytesIO(data)
self._len = len(data)
def read(self, size=None):
"""Read data from the body.
Args:
size: Number of bytes to read, or None for all remaining
Returns:
bytes: The requested data
"""
if size is None:
return self._data.read()
return self._data.read(size)
def readline(self, size=None):
"""Read a line from the body.
Args:
size: Maximum bytes to read
Returns:
bytes: A line of data
"""
if size is None:
return self._data.readline()
return self._data.readline(size)
def readlines(self, hint=None):
"""Read all lines from the body.
Args:
hint: Approximate byte count hint
Returns:
list: List of lines
"""
return self._data.readlines(hint)
def __iter__(self):
"""Iterate over lines in the body."""
return iter(self._data)
def __len__(self):
"""Return the content length."""
return self._len
def close(self):
"""Close the body stream."""
self._data.close()
class HTTP2Request(HeaderPolicy):
"""HTTP/2 request wrapper compatible with gunicorn Request interface.
Wraps an HTTP2Stream to provide the same interface as the HTTP/1.x
Request class, allowing workers to handle HTTP/2 requests using
existing code paths.
"""
#: HTTP/2 carries no 100-continue handshake gunicorn can answer: the
#: response would be written as HTTP/1 bytes onto an HTTP/2 connection.
_policy_expect_continue = False
def __init__(self, stream, cfg, peer_addr):
"""Initialize from an HTTP/2 stream.
Args:
stream: HTTP2Stream instance with received headers/body
cfg: Gunicorn configuration object
peer_addr: Client address tuple (host, port)
"""
self.stream = stream
self.cfg = cfg
self.peer_addr = peer_addr
self.remote_addr = peer_addr
# HTTP/2 version tuple
self.version = (2, 0)
# Parse pseudo-headers
pseudo = stream.get_pseudo_headers()
self.method = pseudo.get(':method', 'GET')
# Derive the scheme from the transport, as HTTP/1 does. A client
# supplied :scheme is honoured only from a peer allowed to speak for
# the connection; otherwise it is ignored rather than rejected, which
# mirrors how an untrusted X-Forwarded-Proto is treated on HTTP/1.
self.scheme = "https" if cfg.is_ssl else "http"
claimed_scheme = pseudo.get(':scheme')
if claimed_scheme and self._peer_is_trusted_proxy():
self.scheme = claimed_scheme
authority = pseudo.get(':authority', '')
path = pseudo.get(':path', '/')
# Parse the path into components
self.uri = path
try:
parts = split_request_uri(path)
self.path = parts.path or ""
self.query = parts.query or ""
self.fragment = parts.fragment or ""
except ValueError:
self.path = path
self.query = ""
self.fragment = ""
# Store authority for Host header equivalent
self._authority = authority
# Convert HTTP/2 headers to HTTP/1.1 style and put them through the
# same policy as HTTP/1, so a rule cannot hold on one protocol and be
# skipped on the other.
self.headers = []
scheme_state = [False]
seen = set()
secure_scheme_headers, forwarder_headers = \
self._peer_trusted_for_forwarded()
for name, value in stream.get_regular_headers():
# Convert to uppercase for WSGI compatibility
name = name.upper()
if RFC9110_5_5_INVALID_AND_DANGEROUS.search(value):
raise InvalidHeader(name, req=self)
kept = self._apply_header_policy(
name, value, scheme_state, seen,
secure_scheme_headers, forwarder_headers,
)
if kept is None:
continue
self.headers.append(kept)
# Set Host header from :authority (RFC 9113 section 8.3.1)
# :authority MUST take precedence over Host header. Runs after the
# policy so a duplicate Host is still rejected rather than replaced.
if authority:
self.headers = [(n, v) for n, v in self.headers if n != 'HOST']
self.headers.append(('HOST', authority))
# Trailers (if any)
self.trailers = []
if stream.trailers:
self.trailers = [
(name.upper(), value)
for name, value in stream.trailers
]
# Body - HTTP/2 streams have complete body data
body_data = stream.get_request_body()
self.body = HTTP2Body(body_data)
# Connection state
self.must_close = False
# Never set on HTTP/2: gunicorn answers it with HTTP/1 bytes written
# straight to the socket, which would corrupt the connection.
self._expected_100_continue = False
# Request numbering (for logging)
self.req_number = stream.stream_id
# HTTP/2 does not use proxy protocol through the data stream
self.proxy_protocol_info = None
# Stream priority (RFC 7540 Section 5.3)
self.priority_weight = stream.priority_weight
self.priority_depends_on = stream.priority_depends_on
def force_close(self):
"""Force the connection to close after this request."""
self.must_close = True
def should_close(self):
"""Check if connection should close after this request.
HTTP/2 connections are persistent by design, but we may still
need to close if explicitly requested.
Returns:
bool: True if connection should close
"""
if self.must_close:
return True
# HTTP/2 connections are persistent, don't close by default
return False
def get_header(self, name):
"""Get a header value by name.
Args:
name: Header name (case-insensitive)
Returns:
str: Header value, or None if not found
"""
name = name.upper()
for h_name, h_value in self.headers:
if h_name == name:
return h_value
return None
@property
def content_length(self):
"""Get the Content-Length header value.
Returns:
int: Content length, or None if not set
"""
cl = self.get_header('CONTENT-LENGTH')
if cl is not None:
try:
return int(cl)
except ValueError:
pass
return None
@property
def content_type(self):
"""Get the Content-Type header value.
Returns:
str: Content type, or None if not set
"""
return self.get_header('CONTENT-TYPE')
def __repr__(self):
return (
f"<HTTP2Request "
f"method={self.method} "
f"path={self.path} "
f"stream_id={self.stream.stream_id}>"
)
__all__ = ['HTTP2Request', 'HTTP2Body']
@@ -0,0 +1,59 @@
#
# This file is part of gunicorn released under the MIT license.
# See the NOTICE for more information.
"""WSGI response writer for HTTP/2 streams."""
from gunicorn.http.wsgi import Response
class HTTP2Response(Response):
"""A WSGI Response that frames its output as HTTP/2 instead of HTTP/1.
Only the wire framing is overridden. Everything the WSGI protocol needs
(``start_response``, header processing, the no-body rules for HEAD, 1xx,
204 and 304, the Content-Length accounting in ``write()``) is inherited,
so HTTP/2 responses obey the same rules as HTTP/1 ones rather than a
parallel set that has to be kept in step by hand.
"""
def __init__(self, req, sock, cfg, h2_conn, stream_id):
# sock is unused: every write goes through the HTTP/2 connection.
# The signature matches Response so wsgi.create() can build either.
super().__init__(req, sock, cfg)
self.h2_conn = h2_conn
self.stream_id = stream_id
self._stream_ended = False
def is_chunked(self):
# HTTP/2 has its own framing; chunked transfer coding is forbidden
# (RFC 9113 section 8.1).
return False
def can_sendfile(self):
# sendfile() writes raw bytes to a socket, which would bypass HTTP/2
# framing entirely. Base Response guards this with cfg.is_ssl, which
# happens to cover HTTP/2 over TLS but not over cleartext.
return False
def send_headers(self):
if self.headers_sent:
return
self.h2_conn.send_response_headers(
self.stream_id, self.status_code, self.headers, end_stream=False
)
self.headers_sent = True
def _emit_body(self, data):
if not data:
return
self.h2_conn.send_data(self.stream_id, data, end_stream=False)
def close(self):
if not self.headers_sent:
self.send_headers()
if self._stream_ended:
return
self._stream_ended = True
trailers = getattr(self, "trailers", None)
self.h2_conn.end_stream(self.stream_id, trailers=trailers)

Some files were not shown because too many files have changed in this diff Show More