Skip to content

Commit 701d72c

Browse files
fix: keep custom extensions registered after worker reload
1 parent 3f43199 commit 701d72c

6 files changed

Lines changed: 58 additions & 15 deletions

File tree

cli.go

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,11 +4,15 @@ package frankenphp
44
import "C"
55
import "unsafe"
66

7+
const cliFailureExitCode = 1
8+
79
// ExecuteScriptCLI executes the PHP script passed as parameter.
810
// It returns the exit status code of the script.
911
func ExecuteScriptCLI(script string, args []string) int {
1012
// Ensure extensions are registered before CLI execution
11-
registerExtensions()
13+
if registerExtensions() != nil {
14+
return cliFailureExitCode
15+
}
1216

1317
cScript := C.CString(script)
1418
defer C.free(unsafe.Pointer(cScript))
@@ -21,7 +25,9 @@ func ExecuteScriptCLI(script string, args []string) int {
2125

2226
func ExecutePHPCode(phpCode string) int {
2327
// Ensure extensions are registered before CLI execution
24-
registerExtensions()
28+
if registerExtensions() != nil {
29+
return cliFailureExitCode
30+
}
2531

2632
cCode := C.CString(phpCode)
2733
defer C.free(unsafe.Pointer(cCode))

ext.go

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -3,27 +3,37 @@ package frankenphp
33
// #include "frankenphp.h"
44
import "C"
55
import (
6+
"errors"
67
"sync"
78
"unsafe"
89
)
910

1011
var (
1112
extensions []*C.zend_module_entry
1213
registerOnce sync.Once
14+
registerErr error
15+
16+
// ErrExtensionRegistration is returned when PHP extension module entries cannot be registered.
17+
ErrExtensionRegistration = errors.New("error registering PHP extensions")
1318
)
1419

1520
// RegisterExtension registers a new PHP extension.
1621
func RegisterExtension(me unsafe.Pointer) {
1722
extensions = append(extensions, (*C.zend_module_entry)(me))
1823
}
1924

20-
func registerExtensions() {
25+
func registerExtensions() error {
2126
if len(extensions) == 0 {
22-
return
27+
return nil
2328
}
2429

2530
registerOnce.Do(func() {
26-
C.register_extensions((**C.zend_module_entry)(unsafe.Pointer(&extensions[0])), C.int(len(extensions)))
31+
if C.register_extensions((**C.zend_module_entry)(unsafe.Pointer(&extensions[0])), C.int(len(extensions))) != C.SUCCESS {
32+
registerErr = ErrExtensionRegistration
33+
return
34+
}
2735
extensions = nil
2836
})
37+
38+
return registerErr
2939
}

frankenphp.c

Lines changed: 14 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
#include <stdint.h>
2727
#include <stdio.h>
2828
#include <stdlib.h>
29+
#include <string.h>
2930
#ifndef PHP_WIN32
3031
#include <unistd.h>
3132
#endif
@@ -1945,17 +1946,25 @@ int register_internal_extensions(void) {
19451946
}
19461947
}
19471948

1948-
modules = NULL;
1949-
modules_len = 0;
1950-
19511949
return SUCCESS;
19521950
}
19531951

1954-
void register_extensions(zend_module_entry **m, int len) {
1955-
modules = m;
1952+
int register_extensions(zend_module_entry **m, int len) {
1953+
zend_module_entry **persistent_modules =
1954+
malloc(sizeof(zend_module_entry *) * len);
1955+
1956+
if (persistent_modules == NULL) {
1957+
return FAILURE;
1958+
}
1959+
1960+
memcpy(persistent_modules, m, sizeof(zend_module_entry *) * len);
1961+
1962+
modules = persistent_modules;
19561963
modules_len = len;
19571964

19581965
original_php_register_internal_extensions_func =
19591966
php_register_internal_extensions_func;
19601967
php_register_internal_extensions_func = register_internal_extensions;
1968+
1969+
return SUCCESS;
19611970
}

frankenphp.go

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -249,7 +249,10 @@ func Init(options ...Option) error {
249249
// Docker/Moby has a similar hack: https://github.com/moby/moby/blob/d828b032a87606ae34267e349bf7f7ccb1f6495a/cmd/dockerd/docker.go#L87-L90
250250
signal.Ignore(syscall.SIGPIPE)
251251

252-
registerExtensions()
252+
if err := registerExtensions(); err != nil {
253+
isRunning = false
254+
return err
255+
}
253256

254257
opt := &opt{}
255258
for _, o := range options {

frankenphp.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -228,6 +228,6 @@ size_t frankenphp_get_thread_memory_usage(uintptr_t thread_index);
228228
void frankenphp_force_kill_thread(force_kill_slot slot);
229229
void frankenphp_release_thread_for_kill(force_kill_slot slot);
230230

231-
void register_extensions(zend_module_entry **m, int len);
231+
int register_extensions(zend_module_entry **m, int len);
232232

233233
#endif

internal/testext/exttest.go

Lines changed: 18 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,11 @@ import (
2222
"github.com/stretchr/testify/require"
2323
)
2424

25+
const (
26+
firstExtensionName = "ext1"
27+
secondExtensionName = "ext2"
28+
)
29+
2530
func testRegisterExtension(t *testing.T) {
2631
frankenphp.RegisterExtension(unsafe.Pointer(&C.module1_entry))
2732
frankenphp.RegisterExtension(unsafe.Pointer(&C.module2_entry))
@@ -30,17 +35,27 @@ func testRegisterExtension(t *testing.T) {
3035
require.Nil(t, err)
3136
defer frankenphp.Shutdown()
3237

38+
assertRegisteredExtensions(t)
39+
40+
frankenphp.RestartWorkers()
41+
42+
assertRegisteredExtensions(t)
43+
}
44+
45+
func assertRegisteredExtensions(t *testing.T) {
46+
t.Helper()
47+
3348
req := httptest.NewRequest("GET", "http://example.com/index.php", nil)
3449
w := httptest.NewRecorder()
3550

36-
req, err = frankenphp.NewRequestWithContext(req, frankenphp.WithRequestDocumentRoot("./testdata", false))
51+
req, err := frankenphp.NewRequestWithContext(req, frankenphp.WithRequestDocumentRoot("./testdata", false))
3752
assert.NoError(t, err)
3853

3954
err = frankenphp.ServeHTTP(w, req)
4055
assert.NoError(t, err)
4156

4257
resp := w.Result()
4358
body, _ := io.ReadAll(resp.Body)
44-
assert.Contains(t, string(body), "ext1")
45-
assert.Contains(t, string(body), "ext2")
59+
assert.Contains(t, string(body), firstExtensionName)
60+
assert.Contains(t, string(body), secondExtensionName)
4661
}

0 commit comments

Comments
 (0)