From 0f040475173943c810612ec60fd28661e27c7a98 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=97=A0=E9=97=BB=E9=A3=8E?= Date: Tue, 16 Jun 2026 19:36:47 +0800 Subject: [PATCH] =?UTF-8?q?=E6=94=AF=E6=8C=81=20Web=20=E7=AB=AF=E5=8F=A3?= =?UTF-8?q?=E5=92=8C=20Sock=20=E5=90=8C=E6=97=B6=E5=90=AF=E7=94=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude --- config.go | 55 ++++++++++++++++++++++++++++---- config_test.go | 86 +++++++++++++++++++++++++++++++++++++++++--------- main.go | 51 +++++++++++++++++++----------- 3 files changed, 152 insertions(+), 40 deletions(-) diff --git a/config.go b/config.go index 98f4a57..3273e95 100644 --- a/config.go +++ b/config.go @@ -52,6 +52,8 @@ type mysqlConfig struct { type webConfig struct { Enabled bool `yaml:"enabled"` + PortEnabled bool `yaml:"port_enabled"` + SocketEnabled bool `yaml:"socket_enabled"` Host string `yaml:"host"` Port int `yaml:"port"` SocketPath string `yaml:"socket_path"` @@ -106,6 +108,8 @@ type rawMySQLConfig struct { type rawWebConfig struct { Enabled *bool `yaml:"enabled"` + PortEnabled *bool `yaml:"port_enabled"` + SocketEnabled *bool `yaml:"socket_enabled"` Host *string `yaml:"host"` Port *int `yaml:"port"` SocketPath *string `yaml:"socket_path"` @@ -143,6 +147,8 @@ func defaultConfig() *config { }, Web: webConfig{ Enabled: true, + PortEnabled: true, + SocketEnabled: defaultWebSocketPath() != "", Host: "0.0.0.0", Port: 8080, SocketPath: defaultWebSocketPath(), @@ -160,12 +166,20 @@ func defaultConfig() *config { // defaultConfigDir 根据操作系统返回配置目录。 func defaultConfigDir() string { - if runtime.GOOS == "windows" { + return defaultConfigDirForGOOS(runtime.GOOS) +} + +func defaultConfigDirForGOOS(goos string) string { + if useRelativeDefaultPath(goos) { return filepath.Join(".", "win", "etc", "mesh_mqtt_go") } return filepath.Join(string(filepath.Separator), "etc", "mesh_mqtt_go") } +func useRelativeDefaultPath(goos string) bool { + return goos == "windows" || goos == "darwin" +} + // defaultConfigPath 返回默认配置文件路径。 func defaultConfigPath() string { return filepath.Join(defaultConfigDir(), configFileName) @@ -184,7 +198,7 @@ func defaultMapTileCacheDir() string { } func defaultMapTileCacheDirForGOOS(goos string) string { - if goos == "windows" { + if useRelativeDefaultPath(goos) { return filepath.Join(".", "win", "srv", "mesh_mqtt_go") } return filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go") @@ -194,19 +208,30 @@ func defaultWebSocketPathForGOOS(goos string) string { if goos == "windows" { return "" } + if useRelativeDefaultPath(goos) { + return filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock") + } return filepath.Join(string(filepath.Separator), "opt", "mesh_mqtt_go", "web.sock") } func clearWebSocketPathOnUnsupportedGOOS(cfg *config, goos string) bool { - if goos != "windows" || cfg.Web.SocketPath == "" { + if goos != "windows" { return false } - cfg.Web.SocketPath = "" - return true + changed := false + if cfg.Web.SocketPath != "" { + cfg.Web.SocketPath = "" + changed = true + } + if cfg.Web.SocketEnabled { + cfg.Web.SocketEnabled = false + changed = true + } + return changed } func defaultSQLitePathForGOOS(goos string) string { - if goos == "windows" { + if useRelativeDefaultPath(goos) { return filepath.Join(".", "win", "etc", "mesh_mqtt_go", "mesh_mqtt_go.db") } return filepath.Join(string(filepath.Separator), "srv", "mesh_mqtt_go", "mesh_mqtt_go.db") @@ -336,6 +361,16 @@ func normalizeConfig(raw rawConfig) (*config, bool) { } else { cfg.Web.Enabled = *raw.Web.Enabled } + if raw.Web.PortEnabled == nil { + changed = true + } else { + cfg.Web.PortEnabled = *raw.Web.PortEnabled + } + if raw.Web.SocketEnabled == nil { + changed = true + } else { + cfg.Web.SocketEnabled = *raw.Web.SocketEnabled + } if raw.Web.Host == nil { changed = true } else { @@ -407,9 +442,15 @@ func validateConfig(cfg *config) error { return fmt.Errorf("invalid database.driver %q: must be sqlite or mysql", cfg.Database.Driver) } if cfg.Web.Enabled { - if cfg.Web.SocketPath == "" && (cfg.Web.Port <= 0 || cfg.Web.Port > 65535) { + if !cfg.Web.PortEnabled && !cfg.Web.SocketEnabled { + return fmt.Errorf("web.port_enabled and web.socket_enabled cannot both be false when web is enabled") + } + if cfg.Web.PortEnabled && (cfg.Web.Port <= 0 || cfg.Web.Port > 65535) { return fmt.Errorf("invalid web port %d: must be 1-65535", cfg.Web.Port) } + if cfg.Web.SocketEnabled && cfg.Web.SocketPath == "" { + return fmt.Errorf("web.socket_path is required when web.socket_enabled is true") + } if cfg.Web.StaticDir == "" { return fmt.Errorf("web.static_dir is required when web is enabled") } diff --git a/config_test.go b/config_test.go index fa86a6e..17da855 100644 --- a/config_test.go +++ b/config_test.go @@ -35,6 +35,13 @@ func TestLoadConfigCreatesDefaultFile(t *testing.T) { if !cfg.Web.Enabled { t.Fatalf("web enabled = false, want true") } + if !cfg.Web.PortEnabled { + t.Fatalf("web port enabled = false, want true") + } + wantSocketEnabled := defaultWebSocketPath() != "" + if cfg.Web.SocketEnabled != wantSocketEnabled { + t.Fatalf("web socket enabled = %t, want %t", cfg.Web.SocketEnabled, wantSocketEnabled) + } if cfg.Web.Port != 8080 { t.Fatalf("web port = %d, want 8080", cfg.Web.Port) } @@ -83,7 +90,7 @@ func TestLoadConfigFillsMissingFields(t *testing.T) { t.Fatal(err) } text := string(data) - for _, want := range []string{"host:", "tls:", "enabled:", "cert_file:", "key_file:", "meshtastic:", "psk:", "database:", "driver:", "sqlite:", "mysql:", "dsn:", "web:", "port:", "socket_path:", "static_dir:", "map_tile_cache_dir:"} { + for _, want := range []string{"host:", "tls:", "enabled:", "cert_file:", "key_file:", "meshtastic:", "psk:", "database:", "driver:", "sqlite:", "mysql:", "dsn:", "web:", "port_enabled:", "socket_enabled:", "port:", "socket_path:", "static_dir:", "map_tile_cache_dir:"} { if !strings.Contains(text, want) { t.Fatalf("completed config missing %q in:\n%s", want, text) } @@ -117,7 +124,7 @@ func TestLoadConfigPreservesExplicitWebFalse(t *testing.T) { if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { t.Fatal(err) } - content := "web:\n enabled: false\n host: 127.0.0.1\n port: 8081\n static_dir: ./public\n" + content := "web:\n enabled: false\n port_enabled: false\n socket_enabled: false\n host: 127.0.0.1\n port: 8081\n static_dir: ./public\n" if err := os.WriteFile(path, []byte(content), 0644); err != nil { t.Fatal(err) } @@ -129,6 +136,9 @@ func TestLoadConfigPreservesExplicitWebFalse(t *testing.T) { if cfg.Web.Enabled { t.Fatalf("web enabled = true, want explicit false") } + if cfg.Web.PortEnabled || cfg.Web.SocketEnabled { + t.Fatalf("web listener enabled = %t/%t, want explicit false/false", cfg.Web.PortEnabled, cfg.Web.SocketEnabled) + } if cfg.Web.Host != "127.0.0.1" || cfg.Web.Port != 8081 || cfg.Web.StaticDir != "./public" { t.Fatalf("web config = %#v", cfg.Web) } @@ -157,11 +167,29 @@ func TestLoadConfigMalformedYAMLDoesNotOverwrite(t *testing.T) { } } +func TestDefaultConfigDirForGOOS(t *testing.T) { + wantRelative := filepath.Join(".", "win", "etc", "mesh_mqtt_go") + for _, goos := range []string{"windows", "darwin"} { + path := defaultConfigDirForGOOS(goos) + if path != wantRelative { + t.Fatalf("%s config dir = %q, want %q", goos, path, wantRelative) + } + } + + linuxPath := defaultConfigDirForGOOS("linux") + wantLinux := filepath.Join(string(filepath.Separator), "etc", "mesh_mqtt_go") + if linuxPath != wantLinux { + t.Fatalf("linux config dir = %q, want %q", linuxPath, wantLinux) + } +} + func TestDefaultMapTileCacheDirForGOOS(t *testing.T) { - windowsPath := defaultMapTileCacheDirForGOOS("windows") - wantWindows := filepath.Join(".", "win", "srv", "mesh_mqtt_go") - if windowsPath != wantWindows { - t.Fatalf("windows map tile cache dir = %q, want %q", windowsPath, wantWindows) + wantRelative := filepath.Join(".", "win", "srv", "mesh_mqtt_go") + for _, goos := range []string{"windows", "darwin"} { + path := defaultMapTileCacheDirForGOOS(goos) + if path != wantRelative { + t.Fatalf("%s map tile cache dir = %q, want %q", goos, path, wantRelative) + } } linuxPath := defaultMapTileCacheDirForGOOS("linux") @@ -176,6 +204,12 @@ func TestDefaultWebSocketPathForGOOS(t *testing.T) { t.Fatalf("windows web socket path = %q, want empty", windowsPath) } + darwinPath := defaultWebSocketPathForGOOS("darwin") + wantDarwin := filepath.Join(".", "win", "opt", "mesh_mqtt_go", "web.sock") + if darwinPath != wantDarwin { + t.Fatalf("darwin web socket path = %q, want %q", darwinPath, wantDarwin) + } + linuxPath := defaultWebSocketPathForGOOS("linux") want := filepath.Join(string(filepath.Separator), "opt", "mesh_mqtt_go", "web.sock") if linuxPath != want { @@ -192,6 +226,9 @@ func TestClearWebSocketPathOnUnsupportedGOOS(t *testing.T) { if cfg.Web.SocketPath != "" { t.Fatalf("windows web socket path = %q, want empty", cfg.Web.SocketPath) } + if cfg.Web.SocketEnabled { + t.Fatalf("windows web socket enabled = true, want false") + } cfg.Web.SocketPath = "/opt/mesh_mqtt_go/web.sock" if clearWebSocketPathOnUnsupportedGOOS(cfg, "linux") { @@ -203,9 +240,12 @@ func TestClearWebSocketPathOnUnsupportedGOOS(t *testing.T) { } func TestDefaultSQLitePathForGOOS(t *testing.T) { - windowsPath := defaultSQLitePathForGOOS("windows") - if !strings.Contains(windowsPath, filepath.Join("win", "etc", "mesh_mqtt_go", "mesh_mqtt_go.db")) { - t.Fatalf("windows sqlite path = %q", windowsPath) + wantRelative := filepath.Join(".", "win", "etc", "mesh_mqtt_go", "mesh_mqtt_go.db") + for _, goos := range []string{"windows", "darwin"} { + path := defaultSQLitePathForGOOS(goos) + if path != wantRelative { + t.Fatalf("%s sqlite path = %q, want %q", goos, path, wantRelative) + } } linuxPath := defaultSQLitePathForGOOS("linux") @@ -238,24 +278,38 @@ func TestValidateConfigDatabase(t *testing.T) { func TestValidateConfigWeb(t *testing.T) { cfg := defaultConfig() - cfg.Web.SocketPath = "" cfg.Web.Port = 0 if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web port") { t.Fatalf("invalid web port error = %v, want web port error", err) } cfg = defaultConfig() - cfg.Web.SocketPath = filepath.Join(string(filepath.Separator), "tmp", "mesh_mqtt_go.sock") + cfg.Web.PortEnabled = false cfg.Web.Port = 0 if err := validateConfig(cfg); err != nil { - t.Fatalf("web socket with invalid port error = %v, want nil", err) + t.Fatalf("disabled web port with invalid port error = %v, want nil", err) } cfg = defaultConfig() + cfg.Web.SocketEnabled = false cfg.Web.SocketPath = "" - cfg.Web.Port = 0 - if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web port") { - t.Fatalf("invalid web port without socket error = %v, want web port error", err) + if err := validateConfig(cfg); err != nil { + t.Fatalf("disabled web socket with empty path error = %v, want nil", err) + } + + cfg = defaultConfig() + cfg.Web.PortEnabled = false + cfg.Web.SocketEnabled = true + cfg.Web.SocketPath = "" + if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web.socket_path") { + t.Fatalf("missing web socket path error = %v, want web.socket_path error", err) + } + + cfg = defaultConfig() + cfg.Web.PortEnabled = false + cfg.Web.SocketEnabled = false + if err := validateConfig(cfg); err == nil || !strings.Contains(err.Error(), "web.port_enabled") { + t.Fatalf("disabled web listeners error = %v, want web.port_enabled error", err) } cfg = defaultConfig() @@ -272,6 +326,8 @@ func TestValidateConfigWeb(t *testing.T) { cfg = defaultConfig() cfg.Web.Enabled = false + cfg.Web.PortEnabled = false + cfg.Web.SocketEnabled = false cfg.Web.Port = 0 cfg.Web.StaticDir = "" if err := validateConfig(cfg); err != nil { diff --git a/main.go b/main.go index 40fffde..b78f9a3 100644 --- a/main.go +++ b/main.go @@ -179,9 +179,11 @@ func parseArgs() (*config, error) { flag.StringVar(&cfg.Database.SQLite.Path, "sqlite-path", cfg.Database.SQLite.Path, "SQLite database file path") flag.StringVar(&cfg.Database.MySQL.DSN, "mysql-dsn", cfg.Database.MySQL.DSN, "MySQL database DSN") flag.BoolVar(&cfg.Web.Enabled, "web", cfg.Web.Enabled, "Enable Gin web server") + flag.BoolVar(&cfg.Web.PortEnabled, "web-port-enabled", cfg.Web.PortEnabled, "Enable web server on TCP host and port") + flag.BoolVar(&cfg.Web.SocketEnabled, "web-socket-enabled", cfg.Web.SocketEnabled, "Enable web server on Unix socket; unsupported on Windows") flag.StringVar(&cfg.Web.Host, "web-host", cfg.Web.Host, "Web server listen host") flag.IntVar(&cfg.Web.Port, "web-port", cfg.Web.Port, "Web server listen port") - flag.StringVar(&cfg.Web.SocketPath, "web-socket-path", cfg.Web.SocketPath, "Web server Unix socket path; empty uses host and port; unsupported on Windows") + flag.StringVar(&cfg.Web.SocketPath, "web-socket-path", cfg.Web.SocketPath, "Web server Unix socket path; unsupported on Windows") flag.StringVar(&cfg.Web.StaticDir, "web-static-dir", cfg.Web.StaticDir, "Web frontend static files directory") flag.StringVar(&cfg.Web.MapTileCacheDir, "web-map-tile-cache-dir", cfg.Web.MapTileCacheDir, "Map tile disk cache root directory") flag.StringVar(&cfg.Web.Admin.Username, "admin-username", cfg.Web.Admin.Username, "Web admin username") @@ -245,31 +247,44 @@ func run(cfg *config) error { } defer forwardManager.StopAll() - var httpServer *http.Server - errCh := make(chan error, 1) + var httpServers []*http.Server + errCh := make(chan error, 2) if cfg.Web.Enabled { sessions, err := newSessionManager(cfg.Web.Admin) if err != nil { return err } mqttStatus := mqttRuntimeStatus{server: server, address: mqttAddr, tls: cfg.MQTT.TLS.Enabled, stats: messageStats, dbQueue: dbQueue} - httpServer = newHTTPServer(cfg.Web, store, sessions, mqttStatus, blocking, forwardManager, settings, botSender) - webAddress := httpServer.Addr - go func() { - if cfg.Web.SocketPath != "" { + handler := newRouter(cfg.Web, store, sessions, mqttStatus, blocking, forwardManager, settings, botSender) + webAddresses := []string{} + if cfg.Web.PortEnabled { + httpServer := &http.Server{ + Addr: net.JoinHostPort(cfg.Web.Host, strconv.Itoa(cfg.Web.Port)), + Handler: handler, + } + httpServers = append(httpServers, httpServer) + webAddresses = append(webAddresses, httpServer.Addr) + go func() { + if err := httpServer.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { + errCh <- err + } + }() + } + if cfg.Web.SocketEnabled { + httpServer := &http.Server{Handler: handler} + httpServers = append(httpServers, httpServer) + webAddresses = append(webAddresses, cfg.Web.SocketPath) + go func() { if err := serveHTTPUnixSocket(httpServer, cfg.Web.SocketPath); err != nil && !errors.Is(err, http.ErrServerClosed) { errCh <- err } - return - } - if err := httpServer.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { - errCh <- err - } - }() - if cfg.Web.SocketPath != "" { - webAddress = cfg.Web.SocketPath + }() } - printJSON(map[string]any{"event": "web_started", "address": webAddress, "static_dir": cfg.Web.StaticDir}) + webStarted := map[string]any{"event": "web_started", "addresses": webAddresses, "static_dir": cfg.Web.StaticDir} + if len(webAddresses) > 0 { + webStarted["address"] = webAddresses[0] + } + printJSON(webStarted) } sigCh := make(chan os.Signal, 1) @@ -281,12 +296,12 @@ func run(cfg *config) error { case runErr = <-errCh: } - if httpServer != nil { + for _, httpServer := range httpServers { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() if err := httpServer.Shutdown(ctx); err != nil && runErr == nil { runErr = err } + cancel() } if err := server.Close(); err != nil && runErr == nil { runErr = err