feat(web): add restart-required state for default model changes (#1499)

- track boot and config default models in gateway status/events
- preserve running, starting, and restarting states during health checks
- add safer gateway restart handling with stronger backend test coverage
- expose restart-required UI and refresh model state after default model update
This commit is contained in:
wenjie 2026-03-13 16:30:59 +08:00 committed by GitHub
parent 4ccea5eb93
commit 87257819f6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 1022 additions and 253 deletions

View file

@ -7,8 +7,11 @@ import (
// GatewayEvent represents a state change event for the gateway process. // GatewayEvent represents a state change event for the gateway process.
type GatewayEvent struct { type GatewayEvent struct {
Status string `json:"gateway_status"` // "running", "starting", "stopped", "error" Status string `json:"gateway_status"` // "running", "starting", "restarting", "stopped", "error"
PID int `json:"pid,omitempty"` PID int `json:"pid,omitempty"`
BootDefaultModel string `json:"boot_default_model,omitempty"`
ConfigDefaultModel string `json:"config_default_model,omitempty"`
RestartRequired bool `json:"gateway_restart_required,omitempty"`
} }
// EventBroadcaster manages SSE client subscriptions and broadcasts events. // EventBroadcaster manages SSE client subscriptions and broadcasts events.

View file

@ -23,13 +23,29 @@ import (
// gateway holds the state for the managed gateway process. // gateway holds the state for the managed gateway process.
var gateway = struct { var gateway = struct {
mu sync.Mutex mu sync.Mutex
cmd *exec.Cmd cmd *exec.Cmd
logs *LogBuffer bootDefaultModel string
events *EventBroadcaster runtimeStatus string
startupDeadline time.Time
logs *LogBuffer
events *EventBroadcaster
}{ }{
logs: NewLogBuffer(200), runtimeStatus: "stopped",
events: NewEventBroadcaster(), logs: NewLogBuffer(200),
events: NewEventBroadcaster(),
}
var (
gatewayStartupWindow = 15 * time.Second
gatewayRestartGracePeriod = 5 * time.Second
gatewayRestartForceKillWindow = 3 * time.Second
gatewayRestartPollInterval = 100 * time.Millisecond
)
var gatewayHealthGet = func(url string, timeout time.Duration) (*http.Response, error) {
client := http.Client{Timeout: timeout}
return client.Get(url)
} }
// registerGatewayRoutes binds gateway lifecycle endpoints to the ServeMux. // registerGatewayRoutes binds gateway lifecycle endpoints to the ServeMux.
@ -65,7 +81,7 @@ func (h *Handler) TryAutoStartGateway() {
return return
} }
pid, err := h.startGatewayLocked() pid, err := h.startGatewayLocked("starting")
if err != nil { if err != nil {
log.Printf("Failed to auto-start gateway: %v", err) log.Printf("Failed to auto-start gateway: %v", err)
return return
@ -131,7 +147,110 @@ func isCmdProcessAliveLocked(cmd *exec.Cmd) bool {
return cmd.Process.Signal(syscall.Signal(0)) == nil return cmd.Process.Signal(syscall.Signal(0)) == nil
} }
func (h *Handler) startGatewayLocked() (int, error) { func setGatewayRuntimeStatusLocked(status string) {
gateway.runtimeStatus = status
if status == "starting" || status == "restarting" {
gateway.startupDeadline = time.Now().Add(gatewayStartupWindow)
return
}
gateway.startupDeadline = time.Time{}
}
func gatewayStatusOnHealthFailureLocked() string {
if gateway.runtimeStatus == "starting" || gateway.runtimeStatus == "restarting" {
if gateway.startupDeadline.IsZero() || time.Now().Before(gateway.startupDeadline) {
return gateway.runtimeStatus
}
return "error"
}
if gateway.runtimeStatus == "running" {
return "running"
}
if gateway.runtimeStatus == "error" {
return "error"
}
return "error"
}
func currentGatewayStatusLocked(processAlive bool) string {
if !processAlive {
if gateway.runtimeStatus == "restarting" {
if gateway.startupDeadline.IsZero() || time.Now().Before(gateway.startupDeadline) {
return "restarting"
}
return "error"
}
if gateway.runtimeStatus == "error" {
return "error"
}
return "stopped"
}
return gatewayStatusOnHealthFailureLocked()
}
func waitForGatewayProcessExit(cmd *exec.Cmd, timeout time.Duration) bool {
if cmd == nil || cmd.Process == nil {
return true
}
deadline := time.Now().Add(timeout)
for {
if !isCmdProcessAliveLocked(cmd) {
return true
}
if time.Now().After(deadline) {
return false
}
time.Sleep(gatewayRestartPollInterval)
}
}
func stopGatewayProcessForRestart(cmd *exec.Cmd) error {
if cmd == nil || cmd.Process == nil || !isCmdProcessAliveLocked(cmd) {
return nil
}
var stopErr error
if runtime.GOOS == "windows" {
stopErr = cmd.Process.Kill()
} else {
stopErr = cmd.Process.Signal(syscall.SIGTERM)
}
if stopErr != nil && isCmdProcessAliveLocked(cmd) {
return fmt.Errorf("failed to stop existing gateway: %w", stopErr)
}
if waitForGatewayProcessExit(cmd, gatewayRestartGracePeriod) {
return nil
}
if runtime.GOOS != "windows" {
killErr := cmd.Process.Signal(syscall.SIGKILL)
if killErr != nil && isCmdProcessAliveLocked(cmd) {
return fmt.Errorf("failed to force-stop existing gateway: %w", killErr)
}
if waitForGatewayProcessExit(cmd, gatewayRestartForceKillWindow) {
return nil
}
}
return fmt.Errorf("existing gateway did not exit before restart")
}
func gatewayRestartRequired(status, bootDefaultModel, configDefaultModel string) bool {
return status == "running" &&
bootDefaultModel != "" &&
configDefaultModel != "" &&
bootDefaultModel != configDefaultModel
}
func (h *Handler) startGatewayLocked(initialStatus string) (int, error) {
cfg, err := config.LoadConfig(h.configPath)
if err != nil {
return 0, fmt.Errorf("failed to load config: %w", err)
}
defaultModelName := strings.TrimSpace(cfg.Agents.Defaults.GetModelName())
// Locate the picoclaw executable // Locate the picoclaw executable
execPath := utils.FindPicoclawBinary() execPath := utils.FindPicoclawBinary()
@ -171,11 +290,19 @@ func (h *Handler) startGatewayLocked() (int, error) {
} }
gateway.cmd = cmd gateway.cmd = cmd
gateway.bootDefaultModel = defaultModelName
setGatewayRuntimeStatusLocked(initialStatus)
pid := cmd.Process.Pid pid := cmd.Process.Pid
log.Printf("Started picoclaw gateway (PID: %d) from %s", pid, execPath) log.Printf("Started picoclaw gateway (PID: %d) from %s", pid, execPath)
// Broadcast starting event // Broadcast the launch state immediately so clients can reflect it without polling.
gateway.events.Broadcast(GatewayEvent{Status: "starting", PID: pid}) gateway.events.Broadcast(GatewayEvent{
Status: initialStatus,
PID: pid,
BootDefaultModel: defaultModelName,
ConfigDefaultModel: defaultModelName,
RestartRequired: false,
})
// Capture stdout/stderr in background // Capture stdout/stderr in background
go scanPipe(stdoutPipe, gateway.logs) go scanPipe(stdoutPipe, gateway.logs)
@ -190,13 +317,23 @@ func (h *Handler) startGatewayLocked() (int, error) {
} }
gateway.mu.Lock() gateway.mu.Lock()
shouldBroadcastStopped := false
if gateway.cmd == cmd { if gateway.cmd == cmd {
gateway.cmd = nil gateway.cmd = nil
gateway.bootDefaultModel = ""
if gateway.runtimeStatus != "restarting" {
setGatewayRuntimeStatusLocked("stopped")
shouldBroadcastStopped = true
}
} }
gateway.mu.Unlock() gateway.mu.Unlock()
// Broadcast stopped event if shouldBroadcastStopped {
gateway.events.Broadcast(GatewayEvent{Status: "stopped"}) gateway.events.Broadcast(GatewayEvent{
Status: "stopped",
RestartRequired: false,
})
}
}() }()
// Start a goroutine to probe health and broadcast "running" once ready // Start a goroutine to probe health and broadcast "running" once ready
@ -219,12 +356,22 @@ func (h *Handler) startGatewayLocked() (int, error) {
healthPort = 18790 healthPort = 18790
} }
healthURL := fmt.Sprintf("http://%s/health", net.JoinHostPort(healthHost, strconv.Itoa(healthPort))) healthURL := fmt.Sprintf("http://%s/health", net.JoinHostPort(healthHost, strconv.Itoa(healthPort)))
client := http.Client{Timeout: 1 * time.Second} resp, err := gatewayHealthGet(healthURL, 1*time.Second)
resp, err := client.Get(healthURL)
if err == nil { if err == nil {
resp.Body.Close() resp.Body.Close()
if resp.StatusCode == http.StatusOK { if resp.StatusCode == http.StatusOK {
gateway.events.Broadcast(GatewayEvent{Status: "running", PID: pid}) gateway.mu.Lock()
if gateway.cmd == cmd {
setGatewayRuntimeStatusLocked("running")
}
gateway.mu.Unlock()
gateway.events.Broadcast(GatewayEvent{
Status: "running",
PID: pid,
BootDefaultModel: defaultModelName,
ConfigDefaultModel: defaultModelName,
RestartRequired: false,
})
return return
} }
} }
@ -253,6 +400,7 @@ func (h *Handler) handleGatewayStart(w http.ResponseWriter, r *http.Request) {
} }
if gateway.cmd != nil && gateway.cmd.Process != nil { if gateway.cmd != nil && gateway.cmd.Process != nil {
gateway.cmd = nil gateway.cmd = nil
setGatewayRuntimeStatusLocked("stopped")
} }
ready, reason, err := h.gatewayStartReady() ready, reason, err := h.gatewayStartReady()
@ -274,7 +422,7 @@ func (h *Handler) handleGatewayStart(w http.ResponseWriter, r *http.Request) {
return return
} }
pid, err := h.startGatewayLocked() pid, err := h.startGatewayLocked("starting")
if err != nil { if err != nil {
http.Error(w, fmt.Sprintf("Failed to start gateway: %v", err), http.StatusInternalServerError) http.Error(w, fmt.Sprintf("Failed to start gateway: %v", err), http.StatusInternalServerError)
return return
@ -330,30 +478,72 @@ func (h *Handler) handleGatewayStop(w http.ResponseWriter, r *http.Request) {
// //
// POST /api/gateway/restart // POST /api/gateway/restart
func (h *Handler) handleGatewayRestart(w http.ResponseWriter, r *http.Request) { func (h *Handler) handleGatewayRestart(w http.ResponseWriter, r *http.Request) {
gateway.mu.Lock() ready, reason, err := h.gatewayStartReady()
if err != nil {
// Stop existing process if running http.Error(
if gateway.cmd != nil && gateway.cmd.Process != nil { w,
if isCmdProcessAliveLocked(gateway.cmd) { fmt.Sprintf("Failed to validate gateway start conditions: %v", err),
// Process is alive, send SIGTERM http.StatusInternalServerError,
if runtime.GOOS == "windows" { )
gateway.cmd.Process.Kill() return
} else { }
gateway.cmd.Process.Signal(syscall.SIGTERM) if !ready {
} w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusBadRequest)
// Wait briefly for it to exit json.NewEncoder(w).Encode(map[string]any{
gateway.mu.Unlock() "status": "precondition_failed",
time.Sleep(2 * time.Second) "message": reason,
gateway.mu.Lock() })
} return
gateway.cmd = nil
} }
gateway.mu.Lock()
previousCmd := gateway.cmd
setGatewayRuntimeStatusLocked("restarting")
gateway.events.Broadcast(GatewayEvent{
Status: "restarting",
RestartRequired: false,
})
gateway.mu.Unlock() gateway.mu.Unlock()
// Start fresh via the existing handler if err = stopGatewayProcessForRestart(previousCmd); err != nil {
h.handleGatewayStart(w, r) gateway.mu.Lock()
if gateway.cmd == previousCmd {
if isCmdProcessAliveLocked(previousCmd) {
setGatewayRuntimeStatusLocked("running")
} else {
gateway.cmd = nil
gateway.bootDefaultModel = ""
setGatewayRuntimeStatusLocked("error")
}
}
gateway.mu.Unlock()
http.Error(w, fmt.Sprintf("Failed to restart gateway: %v", err), http.StatusInternalServerError)
return
}
gateway.mu.Lock()
if gateway.cmd == previousCmd {
gateway.cmd = nil
gateway.bootDefaultModel = ""
}
pid, err := h.startGatewayLocked("restarting")
if err != nil {
gateway.cmd = nil
gateway.bootDefaultModel = ""
setGatewayRuntimeStatusLocked("error")
}
gateway.mu.Unlock()
if err != nil {
http.Error(w, fmt.Sprintf("Failed to restart gateway: %v", err), http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]any{
"status": "ok",
"pid": pid,
})
} }
// handleGatewayClearLogs clears the in-memory gateway log buffer. // handleGatewayClearLogs clears the in-memory gateway log buffer.
@ -374,24 +564,44 @@ func (h *Handler) handleGatewayClearLogs(w http.ResponseWriter, r *http.Request)
// //
// GET /api/gateway/status // GET /api/gateway/status
func (h *Handler) handleGatewayStatus(w http.ResponseWriter, r *http.Request) { func (h *Handler) handleGatewayStatus(w http.ResponseWriter, r *http.Request) {
data := h.gatewayStatusData(r, true)
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(data)
}
func (h *Handler) gatewayStatusData(r *http.Request, includeLogs bool) map[string]any {
data := map[string]any{} data := map[string]any{}
cfg, cfgErr := config.LoadConfig(h.configPath)
configDefaultModel := ""
if cfgErr == nil && cfg != nil {
configDefaultModel = strings.TrimSpace(cfg.Agents.Defaults.GetModelName())
if configDefaultModel != "" {
data["config_default_model"] = configDefaultModel
}
}
// Check process state // Check process state
gateway.mu.Lock() gateway.mu.Lock()
processAlive := isGatewayProcessAliveLocked() processAlive := isGatewayProcessAliveLocked()
bootDefaultModel := ""
if processAlive { if processAlive {
data["pid"] = gateway.cmd.Process.Pid data["pid"] = gateway.cmd.Process.Pid
if gateway.bootDefaultModel != "" {
data["boot_default_model"] = gateway.bootDefaultModel
bootDefaultModel = gateway.bootDefaultModel
}
} }
gateway.mu.Unlock() gateway.mu.Unlock()
if !processAlive { if !processAlive {
data["gateway_status"] = "stopped" gateway.mu.Lock()
data["gateway_status"] = currentGatewayStatusLocked(false)
gateway.mu.Unlock()
} else { } else {
// Process is alive — probe its health endpoint // Process is alive — probe its health endpoint
cfg, err := config.LoadConfig(h.configPath)
host := "127.0.0.1" host := "127.0.0.1"
port := 18790 port := 18790
if err == nil && cfg != nil { if cfgErr == nil && cfg != nil {
host = gatewayProbeHost(h.effectiveGatewayBindHost(cfg)) host = gatewayProbeHost(h.effectiveGatewayBindHost(cfg))
if cfg.Gateway.Port != 0 { if cfg.Gateway.Port != 0 {
port = cfg.Gateway.Port port = cfg.Gateway.Port
@ -399,21 +609,31 @@ func (h *Handler) handleGatewayStatus(w http.ResponseWriter, r *http.Request) {
} }
url := fmt.Sprintf("http://%s/health", net.JoinHostPort(host, strconv.Itoa(port))) url := fmt.Sprintf("http://%s/health", net.JoinHostPort(host, strconv.Itoa(port)))
client := http.Client{Timeout: 2 * time.Second} resp, err := gatewayHealthGet(url, 2*time.Second)
resp, err := client.Get(url)
if err != nil { if err != nil {
data["gateway_status"] = "starting" gateway.mu.Lock()
data["gateway_status"] = currentGatewayStatusLocked(true)
gateway.mu.Unlock()
} else { } else {
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
gateway.mu.Lock()
setGatewayRuntimeStatusLocked("error")
gateway.mu.Unlock()
data["gateway_status"] = "error" data["gateway_status"] = "error"
data["status_code"] = resp.StatusCode data["status_code"] = resp.StatusCode
} else { } else {
var healthData map[string]any var healthData map[string]any
if decErr := json.NewDecoder(resp.Body).Decode(&healthData); decErr != nil { if decErr := json.NewDecoder(resp.Body).Decode(&healthData); decErr != nil {
gateway.mu.Lock()
setGatewayRuntimeStatusLocked("error")
gateway.mu.Unlock()
data["gateway_status"] = "error" data["gateway_status"] = "error"
} else { } else {
gateway.mu.Lock()
setGatewayRuntimeStatusLocked("running")
gateway.mu.Unlock()
for k, v := range healthData { for k, v := range healthData {
data[k] = v data[k] = v
} }
@ -423,6 +643,13 @@ func (h *Handler) handleGatewayStatus(w http.ResponseWriter, r *http.Request) {
} }
} }
status, _ := data["gateway_status"].(string)
data["gateway_restart_required"] = gatewayRestartRequired(
status,
bootDefaultModel,
configDefaultModel,
)
ready, reason, readyErr := h.gatewayStartReady() ready, reason, readyErr := h.gatewayStartReady()
if readyErr != nil { if readyErr != nil {
data["gateway_start_allowed"] = false data["gateway_start_allowed"] = false
@ -434,11 +661,11 @@ func (h *Handler) handleGatewayStatus(w http.ResponseWriter, r *http.Request) {
} }
} }
// Append incremental log data if includeLogs {
appendGatewayLogs(r, data) appendGatewayLogs(r, data)
}
w.Header().Set("Content-Type", "application/json") return data
json.NewEncoder(w).Encode(data)
} }
// appendGatewayLogs reads log_offset and log_run_id query params from the request // appendGatewayLogs reads log_offset and log_run_id query params from the request
@ -524,28 +751,7 @@ func (h *Handler) handleGatewayEvents(w http.ResponseWriter, r *http.Request) {
// currentGatewayStatus returns the current gateway status as a JSON string. // currentGatewayStatus returns the current gateway status as a JSON string.
func (h *Handler) currentGatewayStatus() string { func (h *Handler) currentGatewayStatus() string {
gateway.mu.Lock() data := h.gatewayStatusData(nil, false)
defer gateway.mu.Unlock()
data := map[string]any{
"gateway_status": "stopped",
}
if isGatewayProcessAliveLocked() {
data["gateway_status"] = "running"
data["pid"] = gateway.cmd.Process.Pid
}
ready, reason, readyErr := h.gatewayStartReady()
if readyErr != nil {
data["gateway_start_allowed"] = false
data["gateway_start_reason"] = readyErr.Error()
} else {
data["gateway_start_allowed"] = ready
if !ready {
data["gateway_start_reason"] = reason
}
}
encoded, _ := json.Marshal(data) encoded, _ := json.Marshal(data)
return string(encoded) return string(encoded)
} }

View file

@ -2,19 +2,76 @@ package api
import ( import (
"encoding/json" "encoding/json"
"errors"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"os" "os"
"os/exec"
"path/filepath" "path/filepath"
"runtime"
"strconv" "strconv"
"strings" "strings"
"testing" "testing"
"time"
"github.com/sipeed/picoclaw/pkg/auth" "github.com/sipeed/picoclaw/pkg/auth"
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/web/backend/utils" "github.com/sipeed/picoclaw/web/backend/utils"
) )
func startLongRunningProcess(t *testing.T) *exec.Cmd {
t.Helper()
var cmd *exec.Cmd
if runtime.GOOS == "windows" {
cmd = exec.Command("powershell", "-NoProfile", "-Command", "Start-Sleep -Seconds 30")
} else {
cmd = exec.Command("sleep", "30")
}
if err := cmd.Start(); err != nil {
t.Fatalf("Start() error = %v", err)
}
return cmd
}
func startIgnoringTermProcess(t *testing.T) *exec.Cmd {
t.Helper()
if runtime.GOOS == "windows" {
t.Skip("TERM handling differs on Windows")
}
cmd := exec.Command("sh", "-c", "trap '' TERM; sleep 30")
if err := cmd.Start(); err != nil {
t.Fatalf("Start() error = %v", err)
}
return cmd
}
func resetGatewayTestState(t *testing.T) {
t.Helper()
originalHealthGet := gatewayHealthGet
originalRestartGracePeriod := gatewayRestartGracePeriod
originalRestartForceKillWindow := gatewayRestartForceKillWindow
originalRestartPollInterval := gatewayRestartPollInterval
t.Cleanup(func() {
gatewayHealthGet = originalHealthGet
gatewayRestartGracePeriod = originalRestartGracePeriod
gatewayRestartForceKillWindow = originalRestartForceKillWindow
gatewayRestartPollInterval = originalRestartPollInterval
gateway.mu.Lock()
gateway.cmd = nil
gateway.bootDefaultModel = ""
setGatewayRuntimeStatusLocked("stopped")
gateway.mu.Unlock()
})
}
func TestGatewayStartReady_NoDefaultModel(t *testing.T) { func TestGatewayStartReady_NoDefaultModel(t *testing.T) {
configPath := filepath.Join(t.TempDir(), "config.json") configPath := filepath.Join(t.TempDir(), "config.json")
h := NewHandler(configPath) h := NewHandler(configPath)
@ -317,6 +374,339 @@ func TestGatewayStatusIncludesStartConditionWhenNotReady(t *testing.T) {
} }
} }
func TestGatewayStatusKeepsRunningWhenHealthProbeFailsAfterRunning(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.bootDefaultModel = "existing-model"
// Simulate a process that has already reached the running state.
setGatewayRuntimeStatusLocked("running")
gateway.mu.Unlock()
gatewayHealthGet = func(string, time.Duration) (*http.Response, error) {
return nil, errors.New("probe failed")
}
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 != "running" {
t.Fatalf("gateway_status = %#v, want %q", got, "running")
}
}
func TestGatewayStatusReturnsErrorAfterStartupWindowExpires(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.bootDefaultModel = "existing-model"
setGatewayRuntimeStatusLocked("starting")
gateway.startupDeadline = time.Now().Add(-time.Second)
gateway.mu.Unlock()
gatewayHealthGet = func(string, time.Duration) (*http.Response, error) {
return nil, errors.New("probe failed")
}
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 != "error" {
t.Fatalf("gateway_status = %#v, want %q", got, "error")
}
}
func TestGatewayStatusReturnsRestartingDuringRestartGap(t *testing.T) {
resetGatewayTestState(t)
configPath := filepath.Join(t.TempDir(), "config.json")
h := NewHandler(configPath)
mux := http.NewServeMux()
h.RegisterRoutes(mux)
gateway.mu.Lock()
setGatewayRuntimeStatusLocked("restarting")
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 != "restarting" {
t.Fatalf("gateway_status = %#v, want %q", got, "restarting")
}
}
func TestGatewayStatusIncludesRestartRequiredWhenModelsDiffer(t *testing.T) {
resetGatewayTestState(t)
configPath := filepath.Join(t.TempDir(), "config.json")
cfg := config.DefaultConfig()
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
cfg.ModelList[0].APIKey = "test-key"
if err := config.SaveConfig(configPath, cfg); err != nil {
t.Fatalf("SaveConfig() error = %v", err)
}
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.bootDefaultModel = "previous-model"
setGatewayRuntimeStatusLocked("running")
gateway.mu.Unlock()
gatewayHealthGet = func(string, time.Duration) (*http.Response, error) {
rec := httptest.NewRecorder()
rec.WriteHeader(http.StatusOK)
_, _ = rec.WriteString(`{"ok":true}`)
return rec.Result(), nil
}
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_restart_required"]; got != true {
t.Fatalf("gateway_restart_required = %#v, want true", got)
}
}
func TestGatewayRestartKeepsRunningProcessWhenPreconditionsFail(t *testing.T) {
configPath := filepath.Join(t.TempDir(), "config.json")
cfg := config.DefaultConfig()
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
cfg.ModelList[0].APIKey = ""
cfg.ModelList[0].AuthMethod = ""
if err := config.SaveConfig(configPath, cfg); err != nil {
t.Fatalf("SaveConfig() error = %v", err)
}
h := NewHandler(configPath)
mux := http.NewServeMux()
h.RegisterRoutes(mux)
cmd := startLongRunningProcess(t)
t.Cleanup(func() {
gateway.mu.Lock()
if gateway.cmd == cmd {
gateway.cmd = nil
gateway.bootDefaultModel = ""
}
gateway.mu.Unlock()
if cmd.Process != nil {
_ = cmd.Process.Kill()
}
_ = cmd.Wait()
})
gateway.mu.Lock()
gateway.cmd = cmd
gateway.bootDefaultModel = "existing-model"
gateway.mu.Unlock()
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/gateway/restart", nil)
mux.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest)
}
gateway.mu.Lock()
stillRunning := gateway.cmd == cmd && isCmdProcessAliveLocked(cmd)
gateway.mu.Unlock()
if !stillRunning {
t.Fatalf("gateway process was stopped when restart preconditions failed")
}
}
func TestGatewayRestartKeepsOldProcessWhenItDoesNotExitInTime(t *testing.T) {
resetGatewayTestState(t)
configPath := filepath.Join(t.TempDir(), "config.json")
cfg := config.DefaultConfig()
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
cfg.ModelList[0].APIKey = "test-key"
if err := config.SaveConfig(configPath, cfg); err != nil {
t.Fatalf("SaveConfig() error = %v", err)
}
h := NewHandler(configPath)
mux := http.NewServeMux()
h.RegisterRoutes(mux)
cmd := startIgnoringTermProcess(t)
t.Cleanup(func() {
gateway.mu.Lock()
if gateway.cmd == cmd {
gateway.cmd = nil
gateway.bootDefaultModel = ""
}
gateway.mu.Unlock()
if cmd.Process != nil {
_ = cmd.Process.Kill()
}
_ = cmd.Wait()
})
gatewayRestartGracePeriod = 150 * time.Millisecond
gatewayRestartForceKillWindow = 150 * time.Millisecond
gatewayRestartPollInterval = 10 * time.Millisecond
gateway.mu.Lock()
gateway.cmd = cmd
gateway.bootDefaultModel = "existing-model"
setGatewayRuntimeStatusLocked("running")
gateway.mu.Unlock()
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/gateway/restart", nil)
mux.ServeHTTP(rec, req)
if rec.Code != http.StatusInternalServerError {
t.Fatalf("status = %d, want %d", rec.Code, http.StatusInternalServerError)
}
gateway.mu.Lock()
stillRunning := gateway.cmd == cmd && isCmdProcessAliveLocked(cmd)
status := gateway.runtimeStatus
gateway.mu.Unlock()
if !stillRunning {
t.Fatalf("gateway process was replaced before the old process exited")
}
if status != "running" {
t.Fatalf("runtimeStatus = %q, want %q", status, "running")
}
}
func TestGatewayRestartReturnsErrorStatusWhenReplacementFailsToStart(t *testing.T) {
resetGatewayTestState(t)
configPath := filepath.Join(t.TempDir(), "config.json")
cfg := config.DefaultConfig()
cfg.Agents.Defaults.ModelName = cfg.ModelList[0].ModelName
cfg.ModelList[0].APIKey = "test-key"
if err := config.SaveConfig(configPath, cfg); err != nil {
t.Fatalf("SaveConfig() error = %v", err)
}
invalidBinaryPath := filepath.Join(t.TempDir(), "fake-picoclaw")
if err := os.WriteFile(invalidBinaryPath, []byte("#!/bin/sh\n"), 0o644); err != nil {
t.Fatalf("WriteFile() error = %v", err)
}
t.Setenv("PICOCLAW_BINARY", invalidBinaryPath)
h := NewHandler(configPath)
mux := http.NewServeMux()
h.RegisterRoutes(mux)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/gateway/restart", nil)
mux.ServeHTTP(rec, req)
if rec.Code != http.StatusInternalServerError {
t.Fatalf("restart status = %d, want %d", rec.Code, http.StatusInternalServerError)
}
statusRec := httptest.NewRecorder()
statusReq := httptest.NewRequest(http.MethodGet, "/api/gateway/status", nil)
mux.ServeHTTP(statusRec, statusReq)
if statusRec.Code != http.StatusOK {
t.Fatalf("status = %d, want %d", statusRec.Code, http.StatusOK)
}
var body map[string]any
if err := json.Unmarshal(statusRec.Body.Bytes(), &body); err != nil {
t.Fatalf("unmarshal response: %v", err)
}
if got := body["gateway_status"]; got != "error" {
t.Fatalf("gateway_status = %#v, want %q", got, "error")
}
}
func TestGatewayClearLogsResetsBufferedHistory(t *testing.T) { func TestGatewayClearLogsResetsBufferedHistory(t *testing.T) {
configPath := filepath.Join(t.TempDir(), "config.json") configPath := filepath.Join(t.TempDir(), "config.json")
h := NewHandler(configPath) h := NewHandler(configPath)

View file

@ -1,10 +1,13 @@
// API client for gateway process management. // API client for gateway process management.
interface GatewayStatusResponse { interface GatewayStatusResponse {
gateway_status: "running" | "starting" | "stopped" | "error" gateway_status: "running" | "starting" | "restarting" | "stopped" | "error"
gateway_start_allowed?: boolean gateway_start_allowed?: boolean
gateway_start_reason?: string gateway_start_reason?: string
gateway_restart_required?: boolean
pid?: number pid?: number
boot_default_model?: string
config_default_model?: string
logs?: string[] logs?: string[]
log_total?: number log_total?: number
log_run_id?: number log_run_id?: number

View file

@ -84,7 +84,7 @@ export async function setDefaultModel(
body: JSON.stringify({ model_name: modelName }), body: JSON.stringify({ model_name: modelName }),
}) })
void refreshGatewayState() await refreshGatewayState()
return response return response
} }

View file

@ -6,6 +6,7 @@ import {
IconMoon, IconMoon,
IconPlayerPlay, IconPlayerPlay,
IconPower, IconPower,
IconRefresh,
IconSun, IconSun,
} from "@tabler/icons-react" } from "@tabler/icons-react"
import { Link } from "@tanstack/react-router" import { Link } from "@tanstack/react-router"
@ -31,6 +32,11 @@ import {
} from "@/components/ui/dropdown-menu.tsx" } from "@/components/ui/dropdown-menu.tsx"
import { Separator } from "@/components/ui/separator.tsx" import { Separator } from "@/components/ui/separator.tsx"
import { SidebarTrigger } from "@/components/ui/sidebar" import { SidebarTrigger } from "@/components/ui/sidebar"
import {
Tooltip,
TooltipContent,
TooltipTrigger,
} from "@/components/ui/tooltip"
import { useGateway } from "@/hooks/use-gateway.ts" import { useGateway } from "@/hooks/use-gateway.ts"
import { useTheme } from "@/hooks/use-theme.ts" import { useTheme } from "@/hooks/use-theme.ts"
@ -41,27 +47,35 @@ export function AppHeader() {
state: gwState, state: gwState,
loading: gwLoading, loading: gwLoading,
canStart, canStart,
restartRequired,
start, start,
restart,
stop, stop,
} = useGateway() } = useGateway()
const isRunning = gwState === "running" const isRunning = gwState === "running"
const isStarting = gwState === "starting" const isStarting = gwState === "starting"
const isRestarting = gwState === "restarting"
const isStopped = gwState === "stopped" || gwState === "unknown" const isStopped = gwState === "stopped" || gwState === "unknown"
const showNotConnectedHint = const showNotConnectedHint =
canStart && (gwState === "stopped" || gwState === "error") !isRestarting && canStart && (gwState === "stopped" || gwState === "error")
const [showStopDialog, setShowStopDialog] = React.useState(false) const [showStopDialog, setShowStopDialog] = React.useState(false)
const handleGatewayToggle = () => { const handleGatewayToggle = () => {
if (gwLoading || (!isRunning && !canStart)) return if (gwLoading || isRestarting || (!isRunning && !canStart)) return
if (isRunning) { if (isRunning) {
setShowStopDialog(true) setShowStopDialog(true)
} else { } else {
start() void start()
} }
} }
const handleGatewayRestart = () => {
if (gwLoading || isRestarting || !restartRequired || !canStart) return
void restart()
}
const confirmStop = () => { const confirmStop = () => {
setShowStopDialog(false) setShowStopDialog(false)
stop() stop()
@ -115,35 +129,67 @@ export function AppHeader() {
</AlertDialog> </AlertDialog>
<div className="text-muted-foreground flex items-center gap-1 text-sm font-medium md:gap-2"> <div className="text-muted-foreground flex items-center gap-1 text-sm font-medium md:gap-2">
{restartRequired && (
<Tooltip delayDuration={700}>
<TooltipTrigger asChild>
<Button
variant="secondary"
size="icon-sm"
className="bg-amber-500/15 text-amber-700 hover:bg-amber-500/25 hover:text-amber-800 dark:text-amber-300 dark:hover:bg-amber-500/25"
onClick={handleGatewayRestart}
disabled={gwLoading || isRestarting || !canStart}
aria-label={t("header.gateway.action.restart")}
>
<IconRefresh className="size-4" />
</Button>
</TooltipTrigger>
<TooltipContent>
{t("header.gateway.restartRequired")}
</TooltipContent>
</Tooltip>
)}
{/* Gateway Start/Stop */} {/* Gateway Start/Stop */}
<Button {isRunning ? (
variant={isStarting ? "secondary" : "default"} <Tooltip delayDuration={700}>
size="sm" <TooltipTrigger asChild>
className={`h-8 gap-2 px-3 ${ <Button
isRunning variant="destructive"
? "bg-destructive/10 text-destructive hover:bg-destructive/20" size="icon-sm"
: isStopped className="size-8"
? "bg-green-500 text-white hover:bg-green-600" onClick={handleGatewayToggle}
: "" disabled={gwLoading}
}`} aria-label={t("header.gateway.action.stop")}
onClick={handleGatewayToggle} >
disabled={gwLoading || isStarting || (!isRunning && !canStart)} <IconPower className="h-4 w-4 opacity-80" />
> </Button>
{gwLoading || isStarting ? ( </TooltipTrigger>
<IconLoader2 className="h-4 w-4 animate-spin opacity-70" /> <TooltipContent>{t("header.gateway.action.stop")}</TooltipContent>
) : isRunning ? ( </Tooltip>
<IconPower className="h-4 w-4 opacity-80" /> ) : (
) : ( <Button
<IconPlayerPlay className="h-4 w-4 opacity-80" /> variant={isStarting || isRestarting ? "secondary" : "default"}
)} size="sm"
<span className="text-xs font-semibold"> className={`h-8 gap-2 px-3 ${
{isRunning isStopped ? "bg-green-500 text-white hover:bg-green-600" : ""
? t("header.gateway.action.stop") }`}
: isStarting onClick={handleGatewayToggle}
? t("header.gateway.status.starting") disabled={gwLoading || isStarting || isRestarting || !canStart}
: t("header.gateway.action.start")} >
</span> {gwLoading || isStarting || isRestarting ? (
</Button> <IconLoader2 className="h-4 w-4 animate-spin opacity-70" />
) : (
<IconPlayerPlay className="h-4 w-4 opacity-80" />
)}
<span className="text-xs font-semibold">
{isRestarting
? t("header.gateway.status.restarting")
: isStarting
? t("header.gateway.status.starting")
: t("header.gateway.action.start")}
</span>
</Button>
)}
<Separator <Separator
className="mx-4 my-2 hidden md:block" className="mx-4 my-2 hidden md:block"

View file

@ -20,6 +20,7 @@ export function ChatPage() {
const { t } = useTranslation() const { t } = useTranslation()
const scrollRef = useRef<HTMLDivElement>(null) const scrollRef = useRef<HTMLDivElement>(null)
const [isAtBottom, setIsAtBottom] = useState(true) const [isAtBottom, setIsAtBottom] = useState(true)
const [hasScrolled, setHasScrolled] = useState(false)
const [input, setInput] = useState("") const [input, setInput] = useState("")
const { const {
@ -56,14 +57,22 @@ export function ChatPage() {
onDeletedActiveSession: newChat, onDeletedActiveSession: newChat,
}) })
const handleScroll = (e: React.UIEvent<HTMLDivElement>) => { const syncScrollState = (element: HTMLDivElement) => {
const { scrollTop, scrollHeight, clientHeight } = e.currentTarget const { scrollTop, scrollHeight, clientHeight } = element
setHasScrolled(scrollTop > 0)
setIsAtBottom(scrollHeight - scrollTop <= clientHeight + 10) setIsAtBottom(scrollHeight - scrollTop <= clientHeight + 10)
} }
const handleScroll = (e: React.UIEvent<HTMLDivElement>) => {
syncScrollState(e.currentTarget)
}
useEffect(() => { useEffect(() => {
if (isAtBottom && scrollRef.current) { if (scrollRef.current) {
scrollRef.current.scrollTop = scrollRef.current.scrollHeight if (isAtBottom) {
scrollRef.current.scrollTop = scrollRef.current.scrollHeight
}
syncScrollState(scrollRef.current)
} }
}, [messages, isTyping, isAtBottom]) }, [messages, isTyping, isAtBottom])
@ -77,6 +86,9 @@ export function ChatPage() {
<div className="bg-background/95 flex h-full flex-col"> <div className="bg-background/95 flex h-full flex-col">
<PageHeader <PageHeader
title={t("navigation.chat")} title={t("navigation.chat")}
className={`transition-shadow ${
hasScrolled ? "shadow-sm" : "shadow-none"
}`}
titleExtra={ titleExtra={
hasConfiguredModels && ( hasConfiguredModels && (
<ModelSelector <ModelSelector
@ -90,7 +102,7 @@ export function ChatPage() {
} }
> >
<Button <Button
variant="outline" variant="secondary"
size="sm" size="sm"
onClick={newChat} onClick={newChat}
className="h-9 gap-2" className="h-9 gap-2"

View file

@ -37,7 +37,7 @@ export function ModelSelector({
> >
<SelectValue placeholder={t("chat.noModel")} /> <SelectValue placeholder={t("chat.noModel")} />
</SelectTrigger> </SelectTrigger>
<SelectContent> <SelectContent position="popper" align="start">
{apiKeyModels.length > 0 && ( {apiKeyModels.length > 0 && (
<SelectGroup> <SelectGroup>
<SelectLabel>{t("chat.modelGroup.apikey")}</SelectLabel> <SelectLabel>{t("chat.modelGroup.apikey")}</SelectLabel>

View file

@ -41,7 +41,7 @@ export function SessionHistoryMenu({
return ( return (
<DropdownMenu onOpenChange={onOpenChange}> <DropdownMenu onOpenChange={onOpenChange}>
<DropdownMenuTrigger asChild> <DropdownMenuTrigger asChild>
<Button variant="outline" size="sm" className="h-9 gap-2"> <Button variant="secondary" size="sm" className="h-9 gap-2">
<IconHistory className="size-4" /> <IconHistory className="size-4" />
<span className="hidden sm:inline">{t("chat.history")}</span> <span className="hidden sm:inline">{t("chat.history")}</span>
</Button> </Button>

View file

@ -110,7 +110,7 @@ export function EditModelSheet({
: undefined, : undefined,
thinking_level: form.thinkingLevel || undefined, thinking_level: form.thinkingLevel || undefined,
}) })
if (setAsDefault) { if (setAsDefault && !model.is_default) {
await setDefaultModel(model.model_name) await setDefaultModel(model.model_name)
} }
onSaved() onSaved()

View file

@ -79,6 +79,8 @@ export function ModelsPage() {
}, [fetchModels]) }, [fetchModels])
const handleSetDefault = async (model: ModelInfo) => { const handleSetDefault = async (model: ModelInfo) => {
if (model.is_default) return
setSettingDefaultIndex(model.index) setSettingDefaultIndex(model.index)
try { try {
await setDefaultModel(model.model_name) await setDefaultModel(model.model_name)

View file

@ -2,16 +2,28 @@ import { IconMenu2 } from "@tabler/icons-react"
import type { ReactNode } from "react" import type { ReactNode } from "react"
import { SidebarTrigger } from "@/components/ui/sidebar" import { SidebarTrigger } from "@/components/ui/sidebar"
import { cn } from "@/lib/utils"
interface PageHeaderProps { interface PageHeaderProps {
title: string title: string
titleExtra?: ReactNode titleExtra?: ReactNode
children?: ReactNode children?: ReactNode
className?: string
} }
export function PageHeader({ title, titleExtra, children }: PageHeaderProps) { export function PageHeader({
title,
titleExtra,
children,
className,
}: PageHeaderProps) {
return ( return (
<div className="flex h-14 shrink-0 items-center justify-between px-6 pt-2"> <div
className={cn(
"z-40 flex h-14 shrink-0 items-center justify-between px-6 pt-2",
className,
)}
>
<div className="flex items-center gap-4"> <div className="flex items-center gap-4">
<SidebarTrigger className="border-border/60 bg-background text-muted-foreground hover:bg-accent hover:text-foreground hidden h-9 w-9 rounded-lg border sm:flex [&>svg]:size-5"> <SidebarTrigger className="border-border/60 bg-background text-muted-foreground hover:bg-accent hover:text-foreground hidden h-9 w-9 rounded-lg border sm:flex [&>svg]:size-5">
<IconMenu2 /> <IconMenu2 />

View file

@ -1,4 +1,4 @@
import { useCallback, useEffect, useMemo, useState } from "react" import { useCallback, useEffect, useMemo, useRef, useState } from "react"
import { type ModelInfo, getModels, setDefaultModel } from "@/api/models" import { type ModelInfo, getModels, setDefaultModel } from "@/api/models"
@ -20,6 +20,7 @@ function isLocalModel(model: ModelInfo): boolean {
export function useChatModels({ isConnected }: UseChatModelsOptions) { export function useChatModels({ isConnected }: UseChatModelsOptions) {
const [modelList, setModelList] = useState<ModelInfo[]>([]) const [modelList, setModelList] = useState<ModelInfo[]>([])
const [defaultModelName, setDefaultModelName] = useState("") const [defaultModelName, setDefaultModelName] = useState("")
const setDefaultRequestIdRef = useRef(0)
const loadModels = useCallback(async () => { const loadModels = useCallback(async () => {
try { try {
@ -41,17 +42,28 @@ export function useChatModels({ isConnected }: UseChatModelsOptions) {
return () => clearTimeout(timerId) return () => clearTimeout(timerId)
}, [isConnected, loadModels]) }, [isConnected, loadModels])
const handleSetDefault = useCallback(async (modelName: string) => { const handleSetDefault = useCallback(
try { async (modelName: string) => {
await setDefaultModel(modelName) if (modelName === defaultModelName) return
setDefaultModelName(modelName) const requestId = ++setDefaultRequestIdRef.current
setModelList((prev) =>
prev.map((m) => ({ ...m, is_default: m.model_name === modelName })), try {
) await setDefaultModel(modelName)
} catch (err) { const data = await getModels()
console.error("Failed to set default model:", err) if (requestId !== setDefaultRequestIdRef.current) {
} return
}, []) }
setModelList(data.models)
if (data.models.some((m) => m.model_name === data.default_model)) {
setDefaultModelName(data.default_model)
}
} catch (err) {
console.error("Failed to set default model:", err)
}
},
[defaultModelName],
)
const hasConfiguredModels = useMemo( const hasConfiguredModels = useMemo(
() => modelList.some((m) => m.configured), () => modelList.some((m) => m.configured),

View file

@ -37,7 +37,7 @@ export function useGatewayLogs() {
const fetchLogs = async () => { const fetchLogs = async () => {
if ( if (
!mounted || !mounted ||
(gateway.status !== "running" && gateway.status !== "starting") !["running", "starting", "restarting"].includes(gateway.status)
) { ) {
if (mounted) { if (mounted) {
timeout = setTimeout(fetchLogs, 1000) timeout = setTimeout(fetchLogs, 1000)

View file

@ -1,31 +1,30 @@
import { useAtom } from "jotai" import { useAtomValue } from "jotai"
import { useCallback, useEffect, useState } from "react" import { useCallback, useEffect, useState } from "react"
import { import {
type GatewayStatusResponse, type GatewayStatusResponse,
getGatewayStatus, getGatewayStatus,
restartGateway,
startGateway, startGateway,
stopGateway, stopGateway,
} from "@/api/gateway" } from "@/api/gateway"
import { gatewayAtom } from "@/store" import {
applyGatewayStatusToStore,
gatewayAtom,
updateGatewayStore,
} from "@/store"
// Global variable to ensure we only have one SSE connection // Global variable to ensure we only have one SSE connection
let sseInitialized = false let sseInitialized = false
export function useGateway() { export function useGateway() {
const [{ status: state, canStart }, setGateway] = useAtom(gatewayAtom) const gateway = useAtomValue(gatewayAtom)
const { status: state, canStart, restartRequired } = gateway
const [loading, setLoading] = useState(false) const [loading, setLoading] = useState(false)
const applyGatewayStatus = useCallback( const applyGatewayStatus = useCallback((data: GatewayStatusResponse) => {
(data: GatewayStatusResponse) => { applyGatewayStatusToStore(data)
setGateway((prev) => ({ }, [])
...prev,
status: data.gateway_status ?? "unknown",
canStart: data.gateway_start_allowed ?? true,
}))
},
[setGateway],
)
// Initialize global SSE connection once // Initialize global SSE connection once
useEffect(() => { useEffect(() => {
@ -35,9 +34,10 @@ export function useGateway() {
getGatewayStatus() getGatewayStatus()
.then((data) => applyGatewayStatus(data)) .then((data) => applyGatewayStatus(data))
.catch(() => { .catch(() => {
setGateway({ updateGatewayStore({
status: "unknown", status: "unknown",
canStart: true, canStart: true,
restartRequired: false,
}) })
}) })
@ -59,14 +59,7 @@ export function useGateway() {
data.gateway_status || data.gateway_status ||
typeof data.gateway_start_allowed === "boolean" typeof data.gateway_start_allowed === "boolean"
) { ) {
setGateway((prev) => ({ applyGatewayStatus(data)
...prev,
status: data.gateway_status ?? prev.status,
canStart:
typeof data.gateway_start_allowed === "boolean"
? data.gateway_start_allowed
: prev.canStart,
}))
} }
} catch { } catch {
// ignore // ignore
@ -75,7 +68,9 @@ export function useGateway() {
es.onerror = () => { es.onerror = () => {
// EventSource will auto-reconnect // EventSource will auto-reconnect
setGateway((prev) => ({ ...prev, status: "unknown" })) updateGatewayStore((prev) =>
prev.status === "restarting" ? {} : { status: "unknown" },
)
} }
return () => { return () => {
@ -83,7 +78,7 @@ export function useGateway() {
es.close() es.close()
sseInitialized = false sseInitialized = false
} }
}, [applyGatewayStatus, setGateway]) }, [applyGatewayStatus])
const start = useCallback(async () => { const start = useCallback(async () => {
if (!canStart) return if (!canStart) return
@ -92,19 +87,19 @@ export function useGateway() {
try { try {
await startGateway() await startGateway()
// SSE will push the real state changes, but set optimistic state // SSE will push the real state changes, but set optimistic state
setGateway((prev) => ({ ...prev, status: "starting" })) updateGatewayStore({ status: "starting" })
} catch (err) { } catch (err) {
console.error("Failed to start gateway:", err) console.error("Failed to start gateway:", err)
try { try {
const status = await getGatewayStatus() const status = await getGatewayStatus()
applyGatewayStatus(status) applyGatewayStatus(status)
} catch { } catch {
setGateway((prev) => ({ ...prev, status: "unknown" })) updateGatewayStore({ status: "unknown" })
} }
} finally { } finally {
setLoading(false) setLoading(false)
} }
}, [applyGatewayStatus, canStart, setGateway]) }, [applyGatewayStatus, canStart])
const stop = useCallback(async () => { const stop = useCallback(async () => {
setLoading(true) setLoading(true)
@ -117,5 +112,37 @@ export function useGateway() {
} }
}, []) }, [])
return { state, loading, canStart, start, stop } const restart = useCallback(async () => {
if (state !== "running") return
const previousState = state
const previousCanStart = canStart
const previousRestartRequired = restartRequired
setLoading(true)
updateGatewayStore({
status: "restarting",
restartRequired: false,
})
try {
await restartGateway()
} catch (err) {
console.error("Failed to restart gateway:", err)
try {
const status = await getGatewayStatus()
applyGatewayStatus(status)
} catch {
updateGatewayStore({
status: previousState,
canStart: previousCanStart,
restartRequired: previousRestartRequired,
})
}
} finally {
setLoading(false)
}
}, [applyGatewayStatus, canStart, restartRequired, state])
return { state, loading, canStart, restartRequired, start, stop, restart }
} }

View file

@ -130,8 +130,9 @@ export function usePicoChat() {
const [connectionState, setConnectionState] = const [connectionState, setConnectionState] =
useState<ConnectionState>("disconnected") useState<ConnectionState>("disconnected")
const [isTyping, setIsTyping] = useState(false) const [isTyping, setIsTyping] = useState(false)
const [activeSessionId, setActiveSessionId] = const [activeSessionId, setActiveSessionId] = useState<string>(
useState<string>(() => readStoredSessionId() || generateSessionId()) () => readStoredSessionId() || generateSessionId(),
)
const wsRef = useRef<WebSocket | null>(null) const wsRef = useRef<WebSocket | null>(null)
const isConnectingRef = useRef(false) const isConnectingRef = useRef(false)
@ -144,9 +145,7 @@ export function usePicoChat() {
setMessages((prev) => { setMessages((prev) => {
const next = const next =
typeof nextState === "function" typeof nextState === "function"
? ( ? (nextState as (prevState: ChatMessage[]) => ChatMessage[])(prev)
nextState as (prevState: ChatMessage[]) => ChatMessage[]
)(prev)
: nextState : nextState
if (next !== prev) { if (next !== prev) {
@ -220,64 +219,69 @@ export function usePicoChat() {
} }
}, [loadSessionMessages, setTrackedMessages]) }, [loadSessionMessages, setTrackedMessages])
const handlePicoMessage = useCallback((msg: PicoMessage) => { const handlePicoMessage = useCallback(
const payload = msg.payload || {} (msg: PicoMessage) => {
const payload = msg.payload || {}
switch (msg.type) { switch (msg.type) {
case "message.create": { case "message.create": {
const content = (payload.content as string) || "" const content = (payload.content as string) || ""
const messageId = (payload.message_id as string) || `pico-${Date.now()}` const messageId =
// Use provided timestamp or current time (payload.message_id as string) || `pico-${Date.now()}`
const timestampRaw = // Use provided timestamp or current time
msg.timestamp !== undefined && Number.isFinite(Number(msg.timestamp)) const timestampRaw =
? normalizeUnixTimestamp(Number(msg.timestamp)) msg.timestamp !== undefined &&
: Date.now() Number.isFinite(Number(msg.timestamp))
? normalizeUnixTimestamp(Number(msg.timestamp))
: Date.now()
setTrackedMessages((prev) => [ setTrackedMessages((prev) => [
...prev, ...prev,
{ {
id: messageId, id: messageId,
role: "assistant", role: "assistant",
content, content,
timestamp: timestampRaw, timestamp: timestampRaw,
}, },
]) ])
setIsTyping(false) setIsTyping(false)
break break
}
case "message.update": {
const content = (payload.content as string) || ""
const messageId = payload.message_id as string
if (!messageId) break
setTrackedMessages((prev) =>
prev.map((m) => (m.id === messageId ? { ...m, content } : m)),
)
break
}
case "typing.start":
setIsTyping(true)
break
case "typing.stop":
setIsTyping(false)
break
case "error":
console.error("Pico error:", payload)
setIsTyping(false)
break
case "pong":
// heartbeat response, ignore
break
default:
console.log("Unknown pico message type:", msg.type)
} }
},
case "message.update": { [setTrackedMessages],
const content = (payload.content as string) || "" )
const messageId = payload.message_id as string
if (!messageId) break
setTrackedMessages((prev) =>
prev.map((m) => (m.id === messageId ? { ...m, content } : m)),
)
break
}
case "typing.start":
setIsTyping(true)
break
case "typing.stop":
setIsTyping(false)
break
case "error":
console.error("Pico error:", payload)
setIsTyping(false)
break
case "pong":
// heartbeat response, ignore
break
default:
console.log("Unknown pico message type:", msg.type)
}
}, [setTrackedMessages])
const connect = useCallback(async () => { const connect = useCallback(async () => {
if ( if (
@ -389,32 +393,35 @@ export function usePicoChat() {
return () => disconnect() return () => disconnect()
}, [disconnect]) }, [disconnect])
const sendMessage = useCallback((content: string) => { const sendMessage = useCallback(
if (!wsRef.current || wsRef.current.readyState !== WebSocket.OPEN) { (content: string) => {
console.warn("WebSocket not connected") if (!wsRef.current || wsRef.current.readyState !== WebSocket.OPEN) {
return console.warn("WebSocket not connected")
} return
}
const id = `msg-${++msgIdCounter.current}-${Date.now()}` const id = `msg-${++msgIdCounter.current}-${Date.now()}`
const timestampRaw = Date.now() const timestampRaw = Date.now()
// Add user message to local state // Add user message to local state
setTrackedMessages((prev) => [ setTrackedMessages((prev) => [
...prev, ...prev,
{ id, role: "user", content, timestamp: timestampRaw }, { id, role: "user", content, timestamp: timestampRaw },
]) ])
// Show typing indicator immediately // Show typing indicator immediately
setIsTyping(true) setIsTyping(true)
// Send via Pico Protocol // Send via Pico Protocol
const picoMsg: PicoMessage = { const picoMsg: PicoMessage = {
type: "message.send", type: "message.send",
id, id,
payload: { content }, payload: { content },
} }
wsRef.current.send(JSON.stringify(picoMsg)) wsRef.current.send(JSON.stringify(picoMsg))
}, [setTrackedMessages]) },
[setTrackedMessages],
)
// Switch to a historical session // Switch to a historical session
const switchSession = useCallback( const switchSession = useCallback(
@ -443,7 +450,14 @@ export function usePicoChat() {
} }
}, 100) }, 100)
}, },
[connect, disconnect, gatewayState, loadSessionMessages, setTrackedMessages, t], [
connect,
disconnect,
gatewayState,
loadSessionMessages,
setTrackedMessages,
t,
],
) )
// Start a new empty chat // Start a new empty chat

View file

@ -58,11 +58,14 @@
}, },
"action": { "action": {
"start": "Start Gateway", "start": "Start Gateway",
"stop": "Stop Gateway" "stop": "Stop Gateway",
"restart": "Restart Gateway"
}, },
"status": { "status": {
"starting": "Starting Gateway..." "starting": "Starting Gateway...",
} "restarting": "Restarting Gateway..."
},
"restartRequired": "Model changes require a gateway restart to take effect."
} }
}, },
"common": { "common": {

View file

@ -58,11 +58,14 @@
}, },
"action": { "action": {
"start": "启动服务", "start": "启动服务",
"stop": "停止服务" "stop": "停止服务",
"restart": "重启服务"
}, },
"status": { "status": {
"starting": "服务启动中..." "starting": "服务启动中...",
} "restarting": "服务重启中..."
},
"restartRequired": "切换默认模型后需要重启服务才能生效。"
} }
}, },
"common": { "common": {

View file

@ -5,6 +5,7 @@ import { type GatewayStatusResponse, getGatewayStatus } from "@/api/gateway"
export type GatewayState = export type GatewayState =
| "running" | "running"
| "starting" | "starting"
| "restarting"
| "stopped" | "stopped"
| "error" | "error"
| "unknown" | "unknown"
@ -12,19 +13,54 @@ export type GatewayState =
export interface GatewayStoreState { export interface GatewayStoreState {
status: GatewayState status: GatewayState
canStart: boolean canStart: boolean
restartRequired: boolean
}
type GatewayStorePatch = Partial<GatewayStoreState>
const DEFAULT_GATEWAY_STATE: GatewayStoreState = {
status: "unknown",
canStart: true,
restartRequired: false,
} }
// Global atom for gateway state // Global atom for gateway state
export const gatewayAtom = atom<GatewayStoreState>({ export const gatewayAtom = atom<GatewayStoreState>(DEFAULT_GATEWAY_STATE)
status: "unknown",
canStart: true,
})
function applyGatewayStatusToStore(data: GatewayStatusResponse) { function normalizeGatewayStoreState(
getDefaultStore().set(gatewayAtom, (prev) => ({ prev: GatewayStoreState,
...prev, patch: GatewayStorePatch,
status: data.gateway_status ?? "unknown", ) {
canStart: data.gateway_start_allowed ?? true, return { ...prev, ...patch }
}
export function updateGatewayStore(
patch:
| GatewayStorePatch
| ((prev: GatewayStoreState) => GatewayStorePatch | GatewayStoreState),
) {
getDefaultStore().set(gatewayAtom, (prev) => {
const nextPatch = typeof patch === "function" ? patch(prev) : patch
return normalizeGatewayStoreState(prev, nextPatch)
})
}
export function applyGatewayStatusToStore(
data: Partial<
Pick<
GatewayStatusResponse,
"gateway_status" | "gateway_start_allowed" | "gateway_restart_required"
>
>,
) {
updateGatewayStore((prev) => ({
status: data.gateway_status ?? prev.status,
canStart: data.gateway_start_allowed ?? prev.canStart,
restartRequired:
data.gateway_restart_required ??
(data.gateway_status && data.gateway_status !== "running"
? false
: prev.restartRequired),
})) }))
} }
@ -33,6 +69,6 @@ export async function refreshGatewayState() {
const status = await getGatewayStatus() const status = await getGatewayStatus()
applyGatewayStatusToStore(status) applyGatewayStatusToStore(status)
} catch { } catch {
// Best-effort refresh only; keep current state on error. updateGatewayStore(DEFAULT_GATEWAY_STATE)
} }
} }