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