feat: implement isolated model worker with auto-restart and hot-reload
Tests / Go and GTK tests (push) Successful in 1m38s
Release / Tests before release (push) Successful in 1m34s
Release / Build and publish release (push) Successful in 2m24s

Introduce a supervised child process for capture, VAD, and recognition,
enabling automatic restart on crash/hang and hot-reload via SIGHUP.
Add `--model-auto-restart`, `--model-timeout`, and
`--model-shutdown-timeout` options. Refactor `Run()` with lifecycle
hooks for proper process group cleanup. Update README with new
behavior and CLI flags.
This commit is contained in:
2026-07-17 01:47:37 +03:00
parent b3ed538820
commit b0487be39e
11 changed files with 681 additions and 15 deletions
+36 -6
View File
@@ -14,15 +14,28 @@ import (
"tea.chunkbyte.com/kato/captioneer/src/output"
)
func Run(ctx context.Context, settings config.Settings, sink output.Sink, diagnostics io.Writer) error {
if err := capture.ValidatePrograms(); err != nil {
type directHooks struct {
ready func() error
busy func() error
idle func() error
}
func runDirect(ctx context.Context, settings config.Settings, sink output.Sink, diagnostics io.Writer, hooks directHooks) error {
if err := hooks.busy(); err != nil {
return err
}
resources, err := models.New(settings)
if idleErr := hooks.idle(); err == nil {
err = idleErr
}
if err != nil {
return err
}
defer resources.Close()
defer func() {
_ = hooks.busy()
resources.Close()
_ = hooks.idle()
}()
ctx, cancel := context.WithCancel(ctx)
defer cancel()
@@ -32,24 +45,41 @@ func Run(ctx context.Context, settings config.Settings, sink output.Sink, diagno
return fmt.Errorf("find default output monitor: %w", err)
}
fmt.Fprintf(diagnostics, "Capturing from: %s\n", monitorSource)
fmt.Fprintln(diagnostics, "Press Ctrl+C to stop.")
packets, waitCapture, err := capture.Packets(ctx, monitorSource)
if err != nil {
return fmt.Errorf("start system-audio capture: %w", err)
}
if err := hooks.ready(); err != nil {
cancel()
_ = waitCapture()
return err
}
processor := captions.New(settings, resources, sink.Publish)
var processErr error
for packet := range packets {
if err := processor.Accept(packet); err != nil {
if err := hooks.busy(); err != nil {
processErr = err
cancel()
break
}
acceptErr := processor.Accept(packet)
idleErr := hooks.idle()
if acceptErr != nil || idleErr != nil {
err := errors.Join(acceptErr, idleErr)
processErr = fmt.Errorf("publish caption: %w", err)
cancel()
break
}
}
if processErr == nil {
processErr = processor.Flush()
if err := hooks.busy(); err != nil {
processErr = err
} else {
processErr = processor.Flush()
processErr = errors.Join(processErr, hooks.idle())
}
}
captureErr := waitCapture()
if captureErr != nil && ctx.Err() == nil {
+361
View File
@@ -0,0 +1,361 @@
package app
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"os/exec"
"os/signal"
"strings"
"syscall"
"time"
"tea.chunkbyte.com/kato/captioneer/src/capture"
"tea.chunkbyte.com/kato/captioneer/src/config"
"tea.chunkbyte.com/kato/captioneer/src/models"
"tea.chunkbyte.com/kato/captioneer/src/output"
)
const modelRestartDelay = time.Second
type attemptResult uint8
const (
attemptStopped attemptResult = iota + 1
attemptReload
attemptFailed
attemptFatal
)
type workerHealth struct {
startedAt time.Time
lastSeen time.Time
ready bool
busySince time.Time
}
func (h *workerHealth) observe(messageType string, now time.Time) {
h.lastSeen = now
switch messageType {
case messageReady:
h.ready = true
h.busySince = time.Time{}
case messageBusy:
if h.busySince.IsZero() {
h.busySince = now
}
case messageIdle:
h.busySince = time.Time{}
}
}
func (h workerHealth) timeoutError(now time.Time, timeout time.Duration) error {
if timeout == 0 {
return nil
}
if !h.ready && now.Sub(h.startedAt) >= timeout {
return fmt.Errorf("model worker did not become ready within %s", timeout)
}
if !h.busySince.IsZero() && now.Sub(h.busySince) >= timeout {
return fmt.Errorf("model operation exceeded %s", timeout)
}
if h.ready && now.Sub(h.lastSeen) >= timeout {
return fmt.Errorf("model worker heartbeat was silent for %s", timeout)
}
return nil
}
// Run supervises an isolated process containing capture, VAD, and recognition.
// Presentation stays in this process, so the native model can be replaced
// without restarting the terminal or GTK main loop.
func Run(ctx context.Context, settings config.Settings, sink output.Sink, diagnostics io.Writer) error {
if err := capture.ValidatePrograms(); err != nil {
return err
}
if err := models.Validate(settings); err != nil {
return err
}
fmt.Fprintln(diagnostics, "Press Ctrl+C to stop. Send SIGHUP to hot-reload the model worker.")
reload := make(chan os.Signal, 1)
signal.Notify(reload, syscall.SIGHUP)
defer signal.Stop(reload)
for {
result, err := runWorkerAttempt(ctx, settings, sink, diagnostics, reload)
if ctx.Err() != nil {
return nil
}
switch result {
case attemptStopped:
return nil
case attemptFatal:
return err
case attemptReload:
fmt.Fprintln(diagnostics, "Reloading model worker.")
continue
case attemptFailed:
if !settings.ModelAutoRestart {
return err
}
fmt.Fprintf(diagnostics, "warning: %v; restarting model worker in %s\n", err, modelRestartDelay)
}
timer := time.NewTimer(modelRestartDelay)
select {
case <-ctx.Done():
if !timer.Stop() {
<-timer.C
}
return nil
case <-reload:
if !timer.Stop() {
<-timer.C
}
fmt.Fprintln(diagnostics, "Reloading model worker.")
case <-timer.C:
}
}
}
func runWorkerAttempt(
ctx context.Context,
settings config.Settings,
sink output.Sink,
diagnostics io.Writer,
reload <-chan os.Signal,
) (attemptResult, error) {
cmd, protocol, err := startModelWorker(settings, diagnostics)
if err != nil {
return attemptFailed, fmt.Errorf("start model worker: %w", err)
}
defer protocol.Close()
messages := make(chan workerMessage)
protocolDone := make(chan error, 1)
go readWorkerMessages(protocol, messages, protocolDone)
wait := make(chan error, 1)
go func() { wait <- cmd.Wait() }()
checkInterval := 250 * time.Millisecond
if settings.ModelTimeout > 0 && settings.ModelTimeout < checkInterval {
checkInterval = settings.ModelTimeout / 2
if checkInterval <= 0 {
checkInterval = time.Millisecond
}
}
ticker := time.NewTicker(checkInterval)
defer ticker.Stop()
startedAt := time.Now()
health := workerHealth{startedAt: startedAt, lastSeen: startedAt}
messageChannel := (<-chan workerMessage)(messages)
waitChannel := (<-chan error)(wait)
contextDone := ctx.Done()
var stopTimer *time.Timer
var stopDeadline <-chan time.Time
var stopping bool
var desiredResult attemptResult
var desiredErr error
var processDone bool
var processErr error
var protocolClosed bool
var protocolErr error
var workerFailure string
beginStop := func(result attemptResult, stopErr error) {
if stopping {
if result == attemptFatal {
desiredResult = result
desiredErr = errors.Join(desiredErr, stopErr)
}
return
}
stopping = true
desiredResult = result
desiredErr = stopErr
if err := signalWorkerGroup(cmd, syscall.SIGTERM); err != nil {
desiredErr = errors.Join(desiredErr, err)
}
stopTimer = time.NewTimer(settings.ModelShutdownTimeout)
stopDeadline = stopTimer.C
}
for {
if processDone && protocolClosed {
if stopTimer != nil && !stopTimer.Stop() {
select {
case <-stopTimer.C:
default:
}
}
if stopping {
return desiredResult, desiredErr
}
// The worker may have crashed before its parec child. A process group
// remains signalable after its leader exits, so clean up any survivors
// before starting a replacement worker.
cleanupErr := signalWorkerGroup(cmd, syscall.SIGKILL)
if protocolErr != nil {
return attemptFailed, errors.Join(fmt.Errorf("model worker protocol: %w", protocolErr), cleanupErr)
}
if workerFailure != "" {
return attemptFailed, errors.Join(errors.New(workerFailure), cleanupErr)
}
if processErr != nil {
return attemptFailed, errors.Join(fmt.Errorf("model worker exited: %w", processErr), cleanupErr)
}
return attemptFailed, errors.Join(errors.New("model worker exited unexpectedly"), cleanupErr)
}
select {
case <-contextDone:
contextDone = nil
beginStop(attemptStopped, nil)
case <-reload:
beginStop(attemptReload, nil)
case message, ok := <-messageChannel:
if !ok {
messageChannel = nil
protocolClosed = true
protocolErr = <-protocolDone
if !processDone && !stopping {
beginStop(attemptFailed, errors.New("model worker protocol closed unexpectedly"))
}
continue
}
now := time.Now()
health.observe(message.Type, now)
switch message.Type {
case messageReady:
fmt.Fprintln(diagnostics, "Model worker ready.")
case messageBusy, messageIdle:
case messageHeartbeat:
case messageCaption:
if message.Event == nil {
beginStop(attemptFailed, errors.New("model worker sent an empty caption event"))
continue
}
if err := sink.Publish(*message.Event); err != nil {
beginStop(attemptFatal, fmt.Errorf("publish caption: %w", err))
}
case messageFailure:
workerFailure = message.Error
default:
beginStop(attemptFailed, fmt.Errorf("unknown model worker message %q", message.Type))
}
case err := <-waitChannel:
waitChannel = nil
processDone = true
processErr = err
case now := <-ticker.C:
if !stopping {
if err := health.timeoutError(now, settings.ModelTimeout); err != nil {
beginStop(attemptFailed, err)
}
}
case <-stopDeadline:
stopDeadline = nil
if !processDone {
fmt.Fprintf(diagnostics, "warning: model worker did not stop within %s; force-killing it\n", settings.ModelShutdownTimeout)
if err := signalWorkerGroup(cmd, syscall.SIGKILL); err != nil {
desiredErr = errors.Join(desiredErr, err)
}
}
}
}
}
func startModelWorker(settings config.Settings, diagnostics io.Writer) (*exec.Cmd, *os.File, error) {
executable, err := os.Executable()
if err != nil {
return nil, nil, err
}
encodedSettings, err := encodeWorkerSettings(settings)
if err != nil {
return nil, nil, err
}
protocolRead, protocolWrite, err := os.Pipe()
if err != nil {
return nil, nil, err
}
cmd := exec.Command(executable)
cmd.Env = workerEnvironment(encodedSettings)
cmd.ExtraFiles = []*os.File{protocolWrite}
cmd.Stdout = diagnostics
cmd.Stderr = diagnostics
cmd.SysProcAttr = &syscall.SysProcAttr{
Setpgid: true,
Pdeathsig: syscall.SIGKILL,
}
if err := cmd.Start(); err != nil {
protocolRead.Close()
protocolWrite.Close()
return nil, nil, err
}
if err := protocolWrite.Close(); err != nil {
_ = signalWorkerGroup(cmd, syscall.SIGKILL)
_ = cmd.Wait()
protocolRead.Close()
return nil, nil, err
}
return cmd, protocolRead, nil
}
func workerEnvironment(encodedSettings string) []string {
environment := make([]string, 0, len(os.Environ())+2)
for _, item := range os.Environ() {
if strings.HasPrefix(item, modelWorkerEnvironment+"=") ||
strings.HasPrefix(item, modelSettingsEnvironment+"=") {
continue
}
environment = append(environment, item)
}
return append(environment,
modelWorkerEnvironment+"=1",
modelSettingsEnvironment+"="+encodedSettings,
)
}
func readWorkerMessages(protocol io.Reader, messages chan<- workerMessage, done chan<- error) {
defer close(messages)
decoder := json.NewDecoder(protocol)
for {
var message workerMessage
if err := decoder.Decode(&message); err != nil {
if errors.Is(err, io.EOF) {
done <- nil
} else {
done <- err
}
return
}
messages <- message
}
}
func signalWorkerGroup(cmd *exec.Cmd, signal syscall.Signal) error {
if cmd.Process == nil {
return nil
}
err := syscall.Kill(-cmd.Process.Pid, signal)
if errors.Is(err, syscall.ESRCH) {
return nil
}
if err != nil {
return fmt.Errorf("signal model worker: %w", err)
}
return nil
}
+66
View File
@@ -0,0 +1,66 @@
package app
import (
"strings"
"testing"
"time"
"tea.chunkbyte.com/kato/captioneer/src/config"
)
func TestWorkerHealthDetectsStartupAndOperationTimeouts(t *testing.T) {
started := time.Unix(100, 0)
health := workerHealth{startedAt: started, lastSeen: started}
if err := health.timeoutError(started.Add(time.Second), 2*time.Second); err != nil {
t.Fatalf("healthy startup reported a timeout: %v", err)
}
if err := health.timeoutError(started.Add(2*time.Second), 2*time.Second); err == nil ||
!strings.Contains(err.Error(), "did not become ready") {
t.Fatalf("startup timeout = %v", err)
}
health.observe(messageReady, started.Add(2*time.Second))
health.observe(messageBusy, started.Add(3*time.Second))
if err := health.timeoutError(started.Add(4*time.Second), 2*time.Second); err != nil {
t.Fatalf("healthy operation reported a timeout: %v", err)
}
if err := health.timeoutError(started.Add(5*time.Second), 2*time.Second); err == nil ||
!strings.Contains(err.Error(), "operation exceeded") {
t.Fatalf("operation timeout = %v", err)
}
health.observe(messageIdle, started.Add(5*time.Second))
health.observe(messageHeartbeat, started.Add(6*time.Second))
if err := health.timeoutError(started.Add(7*time.Second), 2*time.Second); err != nil {
t.Fatalf("idle worker reported a timeout: %v", err)
}
if err := health.timeoutError(started.Add(8*time.Second), 2*time.Second); err == nil ||
!strings.Contains(err.Error(), "heartbeat was silent") {
t.Fatalf("heartbeat timeout = %v", err)
}
}
func TestWorkerSettingsRoundTrip(t *testing.T) {
want := config.Settings{
Mode: config.ModeVAD,
ChunkDuration: time.Second,
ModelsDir: "models with spaces",
Threads: 3,
PreviewThresholdDB: -42,
ModelAutoRestart: true,
ModelTimeout: 12 * time.Second,
ModelShutdownTimeout: 2 * time.Second,
}
encoded, err := encodeWorkerSettings(want)
if err != nil {
t.Fatal(err)
}
got, err := decodeWorkerSettings(encoded)
if err != nil {
t.Fatal(err)
}
if got != want {
t.Fatalf("settings round trip = %#v, want %#v", got, want)
}
}
+152
View File
@@ -0,0 +1,152 @@
package app
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"os/signal"
"sync"
"syscall"
"time"
"tea.chunkbyte.com/kato/captioneer/src/captions"
"tea.chunkbyte.com/kato/captioneer/src/config"
)
const (
modelWorkerEnvironment = "CAPTIONEER_INTERNAL_MODEL_WORKER"
modelSettingsEnvironment = "CAPTIONEER_INTERNAL_MODEL_SETTINGS"
modelProtocolFD = 3
)
const (
messageReady = "ready"
messageBusy = "busy"
messageIdle = "idle"
messageCaption = "caption"
messageFailure = "failure"
messageHeartbeat = "heartbeat"
)
type workerMessage struct {
Type string `json:"type"`
Event *captions.Event `json:"event,omitempty"`
Error string `json:"error,omitempty"`
}
type workerReporter struct {
mu sync.Mutex
encoder *json.Encoder
}
func (r *workerReporter) send(message workerMessage) error {
r.mu.Lock()
defer r.mu.Unlock()
return r.encoder.Encode(message)
}
type protocolSink struct {
reporter *workerReporter
}
func (s protocolSink) Publish(event captions.Event) error {
return s.reporter.send(workerMessage{Type: messageCaption, Event: &event})
}
func (protocolSink) Close() error { return nil }
// IsModelWorker reports whether this process was created as Captioneer's
// private inference worker. The environment marker is intentionally internal.
func IsModelWorker() bool {
return os.Getenv(modelWorkerEnvironment) == "1"
}
func encodeWorkerSettings(settings config.Settings) (string, error) {
data, err := json.Marshal(settings)
if err != nil {
return "", err
}
return base64.RawStdEncoding.EncodeToString(data), nil
}
func decodeWorkerSettings(value string) (config.Settings, error) {
data, err := base64.RawStdEncoding.DecodeString(value)
if err != nil {
return config.Settings{}, fmt.Errorf("decode model worker settings: %w", err)
}
var settings config.Settings
if err := json.Unmarshal(data, &settings); err != nil {
return config.Settings{}, fmt.Errorf("parse model worker settings: %w", err)
}
if err := settings.Validate(); err != nil {
return config.Settings{}, fmt.Errorf("validate model worker settings: %w", err)
}
return settings, nil
}
// RunModelWorker owns all Sherpa/VAD resources in the supervised child. It
// returns a process exit code so command entrypoints can call os.Exit.
func RunModelWorker(diagnostics io.Writer) int {
protocol := os.NewFile(uintptr(modelProtocolFD), "captioneer-model-protocol")
if protocol == nil {
fmt.Fprintln(diagnostics, "captioneer: model worker protocol is unavailable")
return 1
}
defer protocol.Close()
syscall.CloseOnExec(modelProtocolFD)
settings, err := decodeWorkerSettings(os.Getenv(modelSettingsEnvironment))
if err != nil {
fmt.Fprintf(diagnostics, "captioneer: %v\n", err)
return 1
}
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
reporter := &workerReporter{encoder: json.NewEncoder(protocol)}
hooks := directHooks{
ready: func() error { return reporter.send(workerMessage{Type: messageReady}) },
busy: func() error { return reporter.send(workerMessage{Type: messageBusy}) },
idle: func() error { return reporter.send(workerMessage{Type: messageIdle}) },
}
heartbeatInterval := time.Second
if settings.ModelTimeout > 0 && settings.ModelTimeout < 4*heartbeatInterval {
heartbeatInterval = settings.ModelTimeout / 4
if heartbeatInterval <= 0 {
heartbeatInterval = time.Millisecond
}
}
heartbeatContext, stopHeartbeat := context.WithCancel(ctx)
heartbeatDone := make(chan struct{})
go func() {
defer close(heartbeatDone)
ticker := time.NewTicker(heartbeatInterval)
defer ticker.Stop()
for {
select {
case <-heartbeatContext.Done():
return
case <-ticker.C:
if reporter.send(workerMessage{Type: messageHeartbeat}) != nil {
return
}
}
}
}()
err = runDirect(ctx, settings, protocolSink{reporter: reporter}, diagnostics, hooks)
stopHeartbeat()
<-heartbeatDone
if err == nil || errors.Is(err, context.Canceled) {
return 0
}
_ = reporter.send(workerMessage{Type: messageFailure, Error: err.Error()})
fmt.Fprintf(diagnostics, "captioneer: model worker: %v\n", err)
return 1
}