actualizado 3-sept

This commit is contained in:
2026-09-03 19:40:44 +02:00
parent 89992823b4
commit 5d2912f789
4338 changed files with 350157 additions and 386 deletions
@@ -0,0 +1,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)