diff --git a/pkg/pid/pidfile.go b/pkg/pid/pidfile.go index 69d02bc65..0b6d461c2 100644 --- a/pkg/pid/pidfile.go +++ b/pkg/pid/pidfile.go @@ -4,6 +4,7 @@ import ( "crypto/rand" "encoding/hex" "encoding/json" + "errors" "fmt" "os" "path/filepath" @@ -16,6 +17,8 @@ import ( const pidFileName = ".picoclaw.pid" +var errInvalidPidFile = errors.New("invalid pid file") + // PidFileData is the JSON structure stored in the PID file. type PidFileData struct { PID int `json:"pid"` @@ -109,6 +112,14 @@ func ReadPidFileWithCheck(homePath string) *PidFileData { pidPath := pidFilePath(homePath) data, err := readPidFileUnlocked(pidPath) if err != nil { + if os.IsNotExist(err) { + return nil + } + if errors.Is(err, errInvalidPidFile) { + logger.Warnf("invalid pid file, remove it: %s (%v)", pidPath, err) + _ = os.Remove(pidPath) + return nil + } logger.Debugf("failed to read pid file: %s", err) return nil } @@ -150,12 +161,12 @@ func readPidFileUnlocked(pidPath string) (*PidFileData, error) { var data PidFileData if err := json.Unmarshal(raw, &data); err != nil { - return nil, err + return nil, fmt.Errorf("%w: %v", errInvalidPidFile, err) } // Validate PID is a positive integer. if data.PID <= 0 { - return nil, fmt.Errorf("invalid pid in pid file: %d", data.PID) + return nil, fmt.Errorf("%w: pid=%d", errInvalidPidFile, data.PID) } return &data, nil diff --git a/pkg/pid/pidfile_test.go b/pkg/pid/pidfile_test.go index 921f590ad..e54b93f4f 100644 --- a/pkg/pid/pidfile_test.go +++ b/pkg/pid/pidfile_test.go @@ -191,6 +191,22 @@ func TestReadPidFileWithCheckStalePID(t *testing.T) { } } +// TestReadPidFileWithCheckInvalidFile auto-cleans malformed PID file. +func TestReadPidFileWithCheckInvalidFile(t *testing.T) { + dir := tmpDir(t) + path := filepath.Join(dir, pidFileName) + os.WriteFile(path, []byte("not json"), 0o600) + + data := ReadPidFileWithCheck(dir) + if data != nil { + t.Error("expected nil for malformed pid file") + } + + if _, err := os.Stat(path); !os.IsNotExist(err) { + t.Error("malformed PID file should be removed") + } +} + // TestRemovePidFile removes the PID file for the current process. func TestRemovePidFile(t *testing.T) { dir := tmpDir(t) diff --git a/pkg/pid/pidfile_unix.go b/pkg/pid/pidfile_unix.go index 5459d8370..7bc53b752 100644 --- a/pkg/pid/pidfile_unix.go +++ b/pkg/pid/pidfile_unix.go @@ -3,6 +3,7 @@ package pid import ( + "errors" "os" "syscall" ) @@ -18,5 +19,11 @@ func isProcessRunning(pid int) bool { return false } // Signal(nil) does not kill the process but checks existence on Unix. - return p.Signal(syscall.Signal(0)) == nil + err = p.Signal(syscall.Signal(0)) + if err == nil { + return true + } + var errno syscall.Errno + // EPERM means the process exists but we are not allowed to signal it. + return errors.As(err, &errno) && errno == syscall.EPERM } diff --git a/web/backend/api/gateway.go b/web/backend/api/gateway.go index b54e55bac..139f2c8c8 100644 --- a/web/backend/api/gateway.go +++ b/web/backend/api/gateway.go @@ -357,7 +357,13 @@ func isCmdProcessAliveLocked(cmd *exec.Cmd) bool { return true } - return cmd.Process.Signal(syscall.Signal(0)) == nil + err := cmd.Process.Signal(syscall.Signal(0)) + if err == nil { + return true + } + var errno syscall.Errno + // EPERM means the process exists but cannot be signaled by this user. + return errors.As(err, &errno) && errno == syscall.EPERM } func setGatewayRuntimeStatusLocked(status string) { @@ -401,6 +407,15 @@ func gatewayStatusWithoutHealthLocked() string { return "error" } if gateway.runtimeStatus == "running" { + // For attached processes there is no waiter goroutine; degrade stale + // running state once the tracked process exits. + if !isCmdProcessAliveLocked(gateway.cmd) { + gateway.cmd = nil + gateway.owned = false + gateway.bootDefaultModel = "" + gateway.bootConfigSignature = "" + return "stopped" + } return "running" } if gateway.runtimeStatus == "error" { @@ -614,6 +629,7 @@ func (h *Handler) startGatewayLocked(initialStatus string, existingPid int) (int // Start a goroutine to probe pidFile and health, update runtime state once ready. go func() { + healthConfirmed := false for i := 0; i < 30; i++ { // try for up to 15 seconds time.Sleep(500 * time.Millisecond) gateway.mu.Lock() @@ -648,7 +664,11 @@ func (h *Handler) startGatewayLocked(initialStatus string, existingPid int) (int setGatewayRuntimeStatusLocked("running") } gateway.mu.Unlock() - return + if !healthConfirmed { + healthConfirmed = true + logger.InfoC("gateway", "Gateway health endpoint reachable; waiting for pid file") + } + continue } } }() @@ -927,8 +947,14 @@ func (h *Handler) gatewayStatusData() map[string]any { // (startGatewayLocked) already handles liveness detection via // pidFile polling and health fallback. gateway.mu.Lock() - data["gateway_status"] = gatewayStatusWithoutHealthLocked() - gateway.pidData = nil + status := gatewayStatusWithoutHealthLocked() + data["gateway_status"] = status + // Keep last known pidData while gateway is still in a transient + // running state; otherwise websocket proxy may lose auth token + // during short pid-file races. + if status == "stopped" || status == "error" { + gateway.pidData = nil + } gateway.mu.Unlock() } diff --git a/web/backend/api/gateway_test.go b/web/backend/api/gateway_test.go index 2ddb1fd8d..1f5f13e27 100644 --- a/web/backend/api/gateway_test.go +++ b/web/backend/api/gateway_test.go @@ -447,6 +447,92 @@ func TestGatewayStatusKeepsRunningWhenHealthProbeFailsAfterRunning(t *testing.T) } } +func TestGatewayStatusKeepsPidDataWhileTrackedProcessAliveWhenPidFileUnavailable(t *testing.T) { + resetGatewayTestState(t) + + configPath := filepath.Join(t.TempDir(), "config.json") + h := NewHandler(configPath) + mux := http.NewServeMux() + h.RegisterRoutes(mux) + + cmd := startLongRunningProcess(t) + t.Cleanup(func() { + if cmd.Process != nil { + _ = cmd.Process.Kill() + } + _ = cmd.Wait() + }) + + gateway.mu.Lock() + gateway.cmd = cmd + gateway.pidData = &ppid.PidFileData{ + PID: cmd.Process.Pid, + Token: "existing-token", + } + setGatewayRuntimeStatusLocked("running") + gateway.mu.Unlock() + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/gateway/status", nil) + mux.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK) + } + + gateway.mu.Lock() + defer gateway.mu.Unlock() + if gateway.pidData == nil { + t.Fatal("gateway.pidData was cleared while runtime status remained running") + } +} + +func TestGatewayStatusDowngradesRunningWhenTrackedProcessExitedAndPidFileMissing(t *testing.T) { + resetGatewayTestState(t) + + configPath := filepath.Join(t.TempDir(), "config.json") + h := NewHandler(configPath) + mux := http.NewServeMux() + h.RegisterRoutes(mux) + + cmd := startLongRunningProcess(t) + if cmd.Process != nil { + _ = cmd.Process.Kill() + } + _ = cmd.Wait() + + gateway.mu.Lock() + gateway.cmd = cmd + gateway.pidData = &ppid.PidFileData{ + PID: cmd.Process.Pid, + Token: "stale-token", + } + setGatewayRuntimeStatusLocked("running") + gateway.mu.Unlock() + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/gateway/status", nil) + mux.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK) + } + + var body map[string]any + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + if got := body["gateway_status"]; got != "stopped" { + t.Fatalf("gateway_status = %#v, want %q", got, "stopped") + } + + gateway.mu.Lock() + defer gateway.mu.Unlock() + if gateway.pidData != nil { + t.Fatal("gateway.pidData should be cleared when tracked process has exited") + } +} + func TestGatewayStatusReportsRunningFromPidProbe(t *testing.T) { resetGatewayTestState(t) diff --git a/web/backend/api/pico.go b/web/backend/api/pico.go index c8ef47308..95bbfd2c1 100644 --- a/web/backend/api/pico.go +++ b/web/backend/api/pico.go @@ -11,6 +11,7 @@ import ( "github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/logger" + ppid "github.com/sipeed/picoclaw/pkg/pid" ) // registerPicoRoutes binds Pico Channel management endpoints to the ServeMux. @@ -57,9 +58,34 @@ func (h *Handler) handleWebSocketProxy() http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { gateway.mu.Lock() ensurePicoTokenCachedLocked(h.configPath) - gatewayAvailable := gateway.pidData != nil + cachedPID := gateway.pidData + trackedCmd := gateway.cmd gateway.mu.Unlock() + gatewayAvailable := false + // Prefer fresh PID file data when available. + if pidData := ppid.ReadPidFileWithCheck(globalConfigDir()); pidData != nil { + gateway.mu.Lock() + gateway.pidData = pidData + setGatewayRuntimeStatusLocked("running") + gatewayAvailable = true + gateway.mu.Unlock() + } else if cachedPID != nil { + // No PID file now: keep availability only while tracked process is + // still alive (covers short PID-file races at startup/restart). + if isCmdProcessAliveLocked(trackedCmd) { + gatewayAvailable = true + } else { + gateway.mu.Lock() + if gateway.cmd == trackedCmd { + gateway.pidData = nil + setGatewayRuntimeStatusLocked("stopped") + } + gatewayAvailable = gateway.pidData != nil + gateway.mu.Unlock() + } + } + if !gatewayAvailable { logger.Warnf("Gateway not available for WebSocket proxy") http.Error(w, "Gateway not available", http.StatusServiceUnavailable) diff --git a/web/backend/api/pico_test.go b/web/backend/api/pico_test.go index ee5586746..04888fde7 100644 --- a/web/backend/api/pico_test.go +++ b/web/backend/api/pico_test.go @@ -11,6 +11,7 @@ import ( "strconv" "testing" + "github.com/sipeed/picoclaw/pkg/channels/pico" "github.com/sipeed/picoclaw/pkg/config" ppid "github.com/sipeed/picoclaw/pkg/pid" ) @@ -307,6 +308,9 @@ func TestHandlePicoSetup_Response(t *testing.T) { } func TestHandleWebSocketProxyReloadsGatewayTargetFromConfig(t *testing.T) { + home := t.TempDir() + t.Setenv("PICOCLAW_HOME", home) + configPath := filepath.Join(t.TempDir(), "config.json") h := NewHandler(configPath) handler := h.handleWebSocketProxy() @@ -335,6 +339,16 @@ func TestHandleWebSocketProxyReloadsGatewayTargetFromConfig(t *testing.T) { if err := config.SaveConfig(configPath, cfg); err != nil { t.Fatalf("SaveConfig() error = %v", err) } + if _, err := ppid.WritePidFile(globalConfigDir(), cfg.Gateway.Host, cfg.Gateway.Port); err != nil { + t.Fatalf("WritePidFile() error = %v", err) + } + origPidData := gateway.pidData + origPicoToken := gateway.picoToken + t.Cleanup(func() { + ppid.RemovePidFile(globalConfigDir()) + gateway.pidData = origPidData + gateway.picoToken = origPicoToken + }) gateway.pidData = &ppid.PidFileData{} gateway.picoToken = "pico" @@ -378,6 +392,9 @@ func TestHandleWebSocketProxyReloadsGatewayTargetFromConfig(t *testing.T) { } func TestHandleWebSocketProxyLoadsCachedPicoTokenWhenMissing(t *testing.T) { + home := t.TempDir() + t.Setenv("PICOCLAW_HOME", home) + configPath := filepath.Join(t.TempDir(), "config.json") h := NewHandler(configPath) handler := h.handleWebSocketProxy() @@ -399,6 +416,12 @@ func TestHandleWebSocketProxyLoadsCachedPicoTokenWhenMissing(t *testing.T) { if err := config.SaveConfig(configPath, cfg); err != nil { t.Fatalf("SaveConfig() error = %v", err) } + if _, err := ppid.WritePidFile(globalConfigDir(), cfg.Gateway.Host, cfg.Gateway.Port); err != nil { + t.Fatalf("WritePidFile() error = %v", err) + } + t.Cleanup(func() { + ppid.RemovePidFile(globalConfigDir()) + }) origPidData := gateway.pidData origPicoToken := gateway.picoToken @@ -426,6 +449,134 @@ func TestHandleWebSocketProxyLoadsCachedPicoTokenWhenMissing(t *testing.T) { } } +func TestHandleWebSocketProxyLoadsPidDataOnDemand(t *testing.T) { + home := t.TempDir() + t.Setenv("PICOCLAW_HOME", home) + + configPath := filepath.Join(t.TempDir(), "config.json") + h := NewHandler(configPath) + handler := h.handleWebSocketProxy() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/pico/ws" { + t.Fatalf("path = %q, want %q", r.URL.Path, "/pico/ws") + } + w.WriteHeader(http.StatusOK) + _, _ = io.WriteString(w, r.Header.Get(protocolKey)) + })) + defer server.Close() + + cfg := config.DefaultConfig() + cfg.Gateway.Host = "127.0.0.1" + cfg.Gateway.Port = mustGatewayTestPort(t, server.URL) + cfg.Channels.Pico.Enabled = true + cfg.Channels.Pico.SetToken("ui-token") + if err := config.SaveConfig(configPath, cfg); err != nil { + t.Fatalf("SaveConfig() error = %v", err) + } + + pidData, err := ppid.WritePidFile(globalConfigDir(), cfg.Gateway.Host, cfg.Gateway.Port) + if err != nil { + t.Fatalf("WritePidFile() error = %v", err) + } + t.Cleanup(func() { + ppid.RemovePidFile(globalConfigDir()) + }) + + origPidData := gateway.pidData + origPicoToken := gateway.picoToken + origStatus := gateway.runtimeStatus + t.Cleanup(func() { + gateway.mu.Lock() + gateway.pidData = origPidData + gateway.picoToken = origPicoToken + gateway.runtimeStatus = origStatus + gateway.mu.Unlock() + }) + + gateway.mu.Lock() + gateway.pidData = nil + gateway.picoToken = "" + setGatewayRuntimeStatusLocked("stopped") + gateway.mu.Unlock() + + req := httptest.NewRequest(http.MethodGet, "/pico/ws?session_id=test-session", nil) + req.Header.Set(protocolKey, tokenPrefix+"ui-token") + rec := httptest.NewRecorder() + handler(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK) + } + + expected := tokenPrefix + pico.PicoTokenPrefix + pidData.Token + "ui-token" + if got := rec.Body.String(); got != expected { + t.Fatalf("forwarded protocol = %q, want %q", got, expected) + } + + gateway.mu.Lock() + defer gateway.mu.Unlock() + if gateway.pidData == nil { + t.Fatal("gateway.pidData should be loaded from pid file") + } + if gateway.runtimeStatus != "running" { + t.Fatalf("runtimeStatus = %q, want %q", gateway.runtimeStatus, "running") + } +} + +func TestHandleWebSocketProxyRejectsStalePidDataAfterProcessExit(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "config.json") + h := NewHandler(configPath) + handler := h.handleWebSocketProxy() + + cfg := config.DefaultConfig() + cfg.Channels.Pico.Enabled = true + cfg.Channels.Pico.SetToken("ui-token") + if err := config.SaveConfig(configPath, cfg); err != nil { + t.Fatalf("SaveConfig() error = %v", err) + } + + cmd := startLongRunningProcess(t) + if cmd.Process != nil { + _ = cmd.Process.Kill() + } + _ = cmd.Wait() + + origPidData := gateway.pidData + origPicoToken := gateway.picoToken + origCmd := gateway.cmd + origStatus := gateway.runtimeStatus + t.Cleanup(func() { + gateway.mu.Lock() + gateway.pidData = origPidData + gateway.picoToken = origPicoToken + gateway.cmd = origCmd + gateway.runtimeStatus = origStatus + gateway.mu.Unlock() + }) + + gateway.mu.Lock() + gateway.pidData = &ppid.PidFileData{PID: cmd.Process.Pid, Token: "stale-token"} + gateway.picoToken = "ui-token" + gateway.cmd = cmd + setGatewayRuntimeStatusLocked("running") + gateway.mu.Unlock() + + req := httptest.NewRequest(http.MethodGet, "/pico/ws?session_id=test-session", nil) + req.Header.Set(protocolKey, tokenPrefix+"ui-token") + rec := httptest.NewRecorder() + handler(rec, req) + + if rec.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusServiceUnavailable) + } + gateway.mu.Lock() + defer gateway.mu.Unlock() + if gateway.pidData != nil { + t.Fatal("gateway.pidData should be cleared after stale process exit is detected") + } +} + func mustGatewayTestPort(t *testing.T, rawURL string) int { t.Helper()