Skip to content

Commit bbb4f33

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

10 files changed

Lines changed: 454 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, and only when the execution environment uses more than one worker. 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: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
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+
Only runs when the execution environment uses more than one worker.
18+
19+
from awslambdaric.lambda_concurrency_hooks import register_pre_fork
20+
21+
@register_pre_fork
22+
def start_inference_server():
23+
subprocess.Popen(["python", "serve.py", "--port", "8000"])
24+
"""
25+
_pre_fork_registry.append((func, (), {}))
26+
return func
27+
28+
29+
def get_pre_fork() -> list[tuple[Callable[..., Any], tuple, dict]]:
30+
return _pre_fork_registry

‎awslambdaric/lambda_multi_concurrent_utils.py‎

Lines changed: 67 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,79 @@ 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+
# One worker behaves like non-concurrent mode, where module-level
106+
# initialization already runs once in that process.
107+
if max_concurrency > 1:
108+
cls._run_pre_fork_hooks(handler, api_addr, log_sink)
109+
110+
# After the hooks, so it is never emitted for a pool that fails to start.
111+
cls._emit_worker_pool_event(max_concurrency)
112+
54113
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).
57114
log_sink.__exit__(None, None, None)
58115

59116
@classmethod
@@ -65,7 +122,7 @@ def run_concurrent(
65122
socket_path: str,
66123
max_concurrency: int,
67124
):
68-
cls._emit_worker_pool_event(max_concurrency)
125+
cls._before_fork(handler, api_addr, max_concurrency)
69126

70127
processes = []
71128
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)