Skip to content

Commit e42e48e

Browse files
committed
feat: Add pre-fork hooks for multi-concurrent mode
1 parent daf3e93 commit e42e48e

10 files changed

Lines changed: 434 additions & 39 deletions

‎RELEASE.CHANGELOG.md‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,7 @@
1+
### September 24, 2026
2+
`4.1.0`
3+
- Add pre-fork hooks for multi-concurrent (Lambda Managed Instances) mode. A function can register callables with the `@register_pre_fork` decorator from `awslambdaric.lambda_concurrency_hooks`; they run once in the parent process, after the handler is imported and before worker processes are started. Workers re-import the handler in their own process, so hooks are for external side effects (starting a subprocess, warming a local service, writing to `/tmp`) and share no in-memory state with workers. A hook that raises is reported to the Runtime API as an INIT error with the type `Runtime.PreForkError` and no worker is started. No impact on the standard on-demand path.
4+
15
### September 15, 2026
26
`4.0.4`
37
- Use the `level` key (instead of `log_level`) for the log level field in JSON-formatted uncaught error logs, aligning it with the key used by other structured log events ([#221](https://github.com/aws/aws-lambda-python-runtime-interface-client/pull/221))

‎awslambdaric/__init__.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,4 +2,4 @@
22
Copyright 2021 Amazon.com, Inc. or its affiliates. All Rights Reserved.
33
"""
44

5-
__version__ = "4.0.4"
5+
__version__ = "4.1.0"

‎awslambdaric/bootstrap.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@
3535
INIT_TYPE_SNAP_START = "snap-start"
3636

3737

38-
def _get_handler(handler):
38+
def get_handler(handler):
3939
try:
4040
modname, fname = handler.rsplit(".", 1)
4141
except ValueError as e:
@@ -516,7 +516,7 @@ def run(handler, lambda_runtime_client):
516516

517517
_log_preview_runtime_warning()
518518

519-
request_handler = _get_handler(handler)
519+
request_handler = get_handler(handler)
520520
except FaultException as e:
521521
error_result = make_error(
522522
e.msg,
Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,31 @@
1+
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2+
# SPDX-License-Identifier: Apache-2.0
3+
4+
from typing import Any, Callable
5+
6+
# The customer-facing name lives only here: the runner consumes get_pre_fork(),
7+
# so a later rename is one line plus a backwards-compatible alias.
8+
__all__ = ["register_pre_fork"]
9+
10+
_pre_fork_registry: list[tuple[Callable[..., Any], tuple, dict]] = []
11+
12+
13+
def register_pre_fork(func: Callable[..., Any]) -> Callable[..., Any]:
14+
"""
15+
Register a function to run once in the parent, before workers are forked.
16+
17+
Runs for any worker count, so a hook's side effects are in place whatever
18+
concurrency the execution environment is configured with.
19+
20+
from awslambdaric.lambda_concurrency_hooks import register_pre_fork
21+
22+
@register_pre_fork
23+
def start_inference_server():
24+
subprocess.Popen(["python", "serve.py", "--port", "8000"])
25+
"""
26+
_pre_fork_registry.append((func, (), {}))
27+
return func
28+
29+
30+
def get_pre_fork() -> list[tuple[Callable[..., Any], tuple, dict]]:
31+
return _pre_fork_registry

‎awslambdaric/lambda_multi_concurrent_utils.py‎

Lines changed: 64 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010

1111
from . import bootstrap
1212
from .lambda_runtime_client import LambdaMultiConcurrentRuntimeClient
13+
from .lambda_concurrency_hooks import get_pre_fork
1314

1415
WORKER_POOL_INITIALIZING_EVENT = "runtime_worker_pool_initializing"
1516

@@ -37,23 +38,76 @@ def run_single(
3738

3839
@classmethod
3940
def _emit_worker_pool_event(cls, max_concurrency: int):
40-
"""Emit worker pool DEBUG event once from the parent before forking.
41-
42-
No output redirection here: RAPID wires the runtime main process's
43-
stdout/stderr to the log egress at spawn. The FD provider socket is
44-
only for the forked workers, which redirect in run_single.
45-
"""
46-
log_sink = bootstrap.init_logging()
41+
"""Emit the worker pool DEBUG event. The sink is owned by _before_fork."""
4742
logging.getLogger().debug(
4843
{
4944
"event": WORKER_POOL_INITIALIZING_EVENT,
5045
"workerCount": max_concurrency,
5146
"executionEnvironmentMaxConcurrency": max_concurrency,
5247
}
5348
)
49+
50+
@classmethod
51+
def _init_handler(cls, handler: str, client, log_sink):
52+
"""Import the handler, mirroring the guard in bootstrap.run: report an
53+
init error to RAPID and exit if it fails."""
54+
try:
55+
return bootstrap.get_handler(handler)
56+
except bootstrap.FaultException as e:
57+
error_result = bootstrap.make_error(e.msg, e.exception_type, e.trace)
58+
except Exception:
59+
error_result = bootstrap.build_fault_result(sys.exc_info(), None)
60+
61+
bootstrap.log_error(error_result, log_sink)
62+
client.post_init_error(error_result)
63+
sys.exit(1)
64+
65+
@classmethod
66+
def _run_pre_fork_hooks(cls, handler: str, api_addr: str, log_sink):
67+
"""Run @register_pre_fork hooks once in the parent, in registration order.
68+
69+
Importing the handler here is what runs its module-level
70+
@register_pre_fork decorators. Workers re-import it in their own
71+
process, so hooks are for external side effects (a subprocess, a warmed
72+
service, a file in /tmp) and share no in-memory state with workers.
73+
74+
A failing hook is reported as an INIT error and exits: no worker should
75+
run against a precondition the hook failed to establish.
76+
"""
77+
client = LambdaMultiConcurrentRuntimeClient(api_addr, False)
78+
cls._init_handler(handler, client, log_sink)
79+
80+
try:
81+
for func, args, kwargs in get_pre_fork():
82+
func(*args, **kwargs)
83+
except Exception:
84+
error_result = bootstrap.build_fault_result(sys.exc_info(), None)
85+
bootstrap.log_error(error_result, log_sink)
86+
client.post_init_error(
87+
error_result, bootstrap.FaultException.PRE_FORK_ERROR
88+
)
89+
sys.exit(1)
90+
91+
@classmethod
92+
def _before_fork(cls, handler: str, api_addr: str, max_concurrency: int):
93+
"""Run the parent's work that must happen before forking workers.
94+
95+
One sink covers both steps, released before returning: forked workers
96+
inherit the parent's handler (fork is the POSIX default before 3.14)
97+
and would log every line twice. Not released on the failure path, where
98+
the process is exiting anyway.
99+
100+
No redirection here: RAPID wires the parent's stdout/stderr to the log
101+
egress at spawn; the FD provider socket is for workers (run_single).
102+
"""
103+
log_sink = bootstrap.init_logging()
104+
105+
cls._run_pre_fork_hooks(handler, api_addr, log_sink)
106+
107+
# After the hooks, so it is never emitted for a pool that fails to start.
108+
cls._emit_worker_pool_event(max_concurrency)
109+
54110
logging.getLogger().handlers.clear()
55-
# Close the sink deterministically now that its handler is gone
56-
# (no-op for StandardLogSink; releases the fd for framed sinks).
57111
log_sink.__exit__(None, None, None)
58112

59113
@classmethod
@@ -65,7 +119,7 @@ def run_concurrent(
65119
socket_path: str,
66120
max_concurrency: int,
67121
):
68-
cls._emit_worker_pool_event(max_concurrency)
122+
cls._before_fork(handler, api_addr, max_concurrency)
69123

70124
processes = []
71125
for _ in range(max_concurrency):

‎awslambdaric/lambda_runtime_exception.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@ class FaultException(Exception):
1313
MALFORMED_HANDLER_NAME = "Runtime.MalformedHandlerName"
1414
BEFORE_SNAPSHOT_ERROR = "Runtime.BeforeSnapshotError"
1515
AFTER_RESTORE_ERROR = "Runtime.AfterRestoreError"
16+
PRE_FORK_ERROR = "Runtime.PreForkError"
1617
LAMBDA_CONTEXT_UNMARSHAL_ERROR = "Runtime.LambdaContextUnmarshalError"
1718
LAMBDA_RUNTIME_CLIENT_ERROR = "Runtime.LambdaRuntimeClientError"
1819

‎tests/test_bootstrap.py‎

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -740,7 +740,7 @@ def __eq__(self, other):
740740
def test_get_event_handler_bad_handler(self):
741741
handler_name = "bad_handler"
742742
with self.assertRaises(FaultException) as cm:
743-
response_handler = bootstrap._get_handler(handler_name)
743+
response_handler = bootstrap.get_handler(handler_name)
744744
returned_exception = cm.exception
745745
self.assertEqual(
746746
self.FaultExceptionMatcher(
@@ -753,7 +753,7 @@ def test_get_event_handler_bad_handler(self):
753753
def test_get_event_handler_import_error(self):
754754
handler_name = "no_module.handler"
755755
with self.assertRaises(FaultException) as cm:
756-
response_handler = bootstrap._get_handler(handler_name)
756+
response_handler = bootstrap.get_handler(handler_name)
757757
returned_exception = cm.exception
758758
self.assertEqual(
759759
self.FaultExceptionMatcher(
@@ -778,7 +778,7 @@ def test_get_event_handler_syntax_error(self):
778778
handler_name = "{}.syntax_error".format(filename)
779779

780780
with self.assertRaises(FaultException) as cm:
781-
response_handler = bootstrap._get_handler(handler_name)
781+
response_handler = bootstrap.get_handler(handler_name)
782782
returned_exception = cm.exception
783783
self.assertEqual(
784784
self.FaultExceptionMatcher(
@@ -801,7 +801,7 @@ def test_get_event_handler_missing_error(self):
801801
filename, _ = os.path.splitext(filename_w_ext)
802802
handler_name = "{}.my_handler".format(filename)
803803
with self.assertRaises(FaultException) as cm:
804-
response_handler = bootstrap._get_handler(handler_name)
804+
response_handler = bootstrap.get_handler(handler_name)
805805
returned_exception = cm.exception
806806
self.assertEqual(
807807
self.FaultExceptionMatcher(
@@ -814,12 +814,12 @@ def test_get_event_handler_missing_error(self):
814814
def test_get_event_handler_slash(self):
815815
importlib.invalidate_caches()
816816
handler_name = "tests/test_handler_with_slash/test_handler.my_handler"
817-
response_handler = bootstrap._get_handler(handler_name)
817+
response_handler = bootstrap.get_handler(handler_name)
818818
response_handler()
819819

820820
def test_get_event_handler_build_in_conflict(self):
821821
with self.assertRaises(FaultException) as cm:
822-
response_handler = bootstrap._get_handler("sys.hello")
822+
response_handler = bootstrap.get_handler("sys.hello")
823823
returned_exception = cm.exception
824824
self.assertEqual(
825825
self.FaultExceptionMatcher(
@@ -830,13 +830,13 @@ def test_get_event_handler_build_in_conflict(self):
830830
)
831831

832832
def test_get_event_handler_doesnt_throw_build_in_module_name_slash(self):
833-
response_handler = bootstrap._get_handler(
833+
response_handler = bootstrap.get_handler(
834834
"tests/test_built_in_module_name/sys.my_handler"
835835
)
836836
response_handler()
837837

838838
def test_get_event_handler_doent_throw_build_in_module_name(self):
839-
response_handler = bootstrap._get_handler(
839+
response_handler = bootstrap.get_handler(
840840
"tests.test_built_in_module_name.sys.my_handler"
841841
)
842842
response_handler()

‎tests/test_concurrency.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,7 @@ def fake_bootstrap_run(handler, lambda_runtime_client):
3939
with patch(
4040
"awslambdaric.lambda_multi_concurrent_utils.MultiConcurrentRunner._redirect_output"
4141
), patch(
42-
"awslambdaric.lambda_multi_concurrent_utils.MultiConcurrentRunner._emit_worker_pool_event"
42+
"awslambdaric.lambda_multi_concurrent_utils.MultiConcurrentRunner._before_fork"
4343
), patch(
4444
"awslambdaric.lambda_multi_concurrent_utils.bootstrap.run",
4545
side_effect=fake_bootstrap_run,

0 commit comments

Comments
 (0)