Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 42 additions & 3 deletions backend/cmd/server/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,14 @@ import (
_ "embed"
"errors"
"flag"
"fmt"
"log"
"net"
"net/http"
_ "net/http/pprof" //nolint:gosec // Admin/debug profiling is intentionally exposed only when the server is started with that route mounted.
"os"
"os/signal"
"path/filepath"
"strconv"
"strings"
"syscall"
Expand Down Expand Up @@ -125,7 +128,10 @@ func runSetupServer() {
IdleTimeout: 120 * time.Second,
}

if err := server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
if err := serveServer(server, config.ServerListenSpec{
Network: config.ServerListenNetworkTCP,
Address: addr,
}); err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Fatalf("Failed to start setup server: %v", err)
}
}
Expand Down Expand Up @@ -154,15 +160,19 @@ func runMainServer() {
defer app.Cleanup()

pprofServer := startPprofServer()
listenSpec, err := cfg.Server.ListenSpec()
if err != nil {
log.Fatalf("Invalid server listen configuration: %v", err)
}

// 启动服务器
go func() {
if err := app.Server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
if err := serveServer(app.Server, listenSpec); err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Fatalf("Failed to start server: %v", err)
}
}()

log.Printf("Server started on %s", app.Server.Addr)
log.Printf("Server started on %s", listenSpec.DisplayAddress())

// 等待中断信号
quit := make(chan os.Signal, 1)
Expand All @@ -187,6 +197,35 @@ func runMainServer() {
log.Println("Server exited")
}

func serveServer(server *http.Server, spec config.ServerListenSpec) error {
switch spec.Network {
case config.ServerListenNetworkUnix:
if err := os.MkdirAll(filepath.Dir(spec.Address), 0o755); err != nil {
return fmt.Errorf("create unix socket directory: %w", err)
}
if err := config.RemoveUnixSocketIfExists(spec.Address); err != nil {
return err
}
listener, err := net.Listen(string(spec.Network), spec.Address)
if err != nil {
return fmt.Errorf("listen on %s: %w", spec.DisplayAddress(), err)
}
if err := os.Chmod(spec.Address, spec.Mode); err != nil {
_ = listener.Close()
_ = os.Remove(spec.Address)
return fmt.Errorf("chmod unix socket %s: %w", spec.Address, err)
}
defer func() {
_ = listener.Close()
_ = os.Remove(spec.Address)
}()
return server.Serve(listener)
default:
server.Addr = spec.Address
return server.ListenAndServe()
}
}

func startPprofServer() *http.Server {
enabledValue := strings.TrimSpace(os.Getenv("PPROF_ENABLED"))
if enabledValue == "" {
Expand Down
223 changes: 223 additions & 0 deletions backend/internal/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,13 @@ package config
import (
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"net/url"
"os"
pathpkg "path"
"strconv"
"strings"
"time"

Expand Down Expand Up @@ -530,6 +533,139 @@ type ServerConfig struct {
H2C H2CConfig `mapstructure:"h2c"` // HTTP/2 Cleartext 配置
}

type ServerListenNetwork string

const (
ServerListenNetworkTCP ServerListenNetwork = "tcp"
ServerListenNetworkUnix ServerListenNetwork = "unix"
)

const (
defaultUnixSocketFileMode os.FileMode = 0o660
)

type ServerListenSpec struct {
Network ServerListenNetwork
Address string
Mode os.FileMode
}

func (s ServerListenSpec) DisplayAddress() string {
switch s.Network {
case ServerListenNetworkUnix:
return fmt.Sprintf("unix://%s", s.Address)
default:
return s.Address
}
}

func ParseServerListenSpec(host string, port int) (ServerListenSpec, error) {
rawHost := strings.TrimSpace(host)
if rawHost == "" {
rawHost = "0.0.0.0"
}

socketMode := defaultUnixSocketFileMode
switch {
case strings.HasPrefix(rawHost, "unix:"):
return parseUnixSocketListenSpec(strings.TrimSpace(strings.TrimPrefix(rawHost, "unix:")), socketMode)
case strings.HasPrefix(rawHost, "/"):
return parseUnixSocketListenSpec(rawHost, socketMode)
default:
if port <= 0 || port > 65535 {
return ServerListenSpec{}, fmt.Errorf("tcp listen port must be between 1-65535")
}
return ServerListenSpec{
Network: ServerListenNetworkTCP,
Address: fmt.Sprintf("%s:%d", rawHost, port),
}, nil
}
}

func parseUnixSocketListenSpec(raw string, defaultMode os.FileMode) (ServerListenSpec, error) {
socketPath := strings.TrimSpace(raw)
if socketPath == "" {
return ServerListenSpec{}, fmt.Errorf("unix socket path is required")
}

mode := defaultMode
if idx := strings.LastIndex(socketPath, ","); idx >= 0 {
maybeMode := strings.TrimSpace(socketPath[idx+1:])
if maybeMode != "" {
parsedMode, err := parseUnixSocketFileMode(maybeMode)
if err != nil {
return ServerListenSpec{}, err
}
mode = parsedMode
socketPath = strings.TrimSpace(socketPath[:idx])
}
}

if socketPath == "" {
return ServerListenSpec{}, fmt.Errorf("unix socket path is required")
}
if !strings.HasPrefix(socketPath, "/") {
return ServerListenSpec{}, fmt.Errorf("unix socket path must be absolute")
}
cleanPath := pathpkg.Clean(socketPath)
if cleanPath == "/" {
return ServerListenSpec{}, fmt.Errorf("unix socket path cannot be root directory")
}

return ServerListenSpec{
Network: ServerListenNetworkUnix,
Address: cleanPath,
Mode: mode,
}, nil
}

func parseUnixSocketFileMode(raw string) (os.FileMode, error) {
value := strings.TrimSpace(raw)
if value == "" {
return 0, fmt.Errorf("unix socket file mode cannot be empty")
}
if strings.HasPrefix(value, "0o") || strings.HasPrefix(value, "0O") {
value = value[2:]
}
if strings.HasPrefix(value, "0") && len(value) > 1 {
value = value[1:]
}
if len(value) < 3 || len(value) > 4 {
return 0, fmt.Errorf("unix socket file mode must be 3-4 octal digits")
}
for _, r := range value {
if r < '0' || r > '7' {
return 0, fmt.Errorf("unix socket file mode must be octal")
}
}
parsed, err := strconv.ParseUint(value, 8, 32)
if err != nil {
return 0, fmt.Errorf("parse unix socket file mode: %w", err)
}
mode := os.FileMode(parsed)
if mode&os.ModeType != 0 {
return 0, fmt.Errorf("unix socket file mode must not include file type bits")
}
return mode, nil
}

func RemoveUnixSocketIfExists(path string) error {
info, err := os.Lstat(path)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil
}
return fmt.Errorf("stat unix socket %q: %w", path, err)
}
if info.Mode()&os.ModeSocket == 0 {
return fmt.Errorf("refusing to remove non-socket file at %q", path)
}
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("remove stale unix socket %q: %w", path, err)
}
return nil
}

// H2CConfig HTTP/2 Cleartext 配置
type H2CConfig struct {
Enabled bool `mapstructure:"enabled"` // 是否启用 H2C
Expand Down Expand Up @@ -967,6 +1103,21 @@ func (s *ServerConfig) Address() string {
return fmt.Sprintf("%s:%d", s.Host, s.Port)
}

func (s *ServerConfig) ListenSpec() (ServerListenSpec, error) {
if s == nil {
return ParseServerListenSpec("", 8080)
}
return ParseServerListenSpec(s.Host, s.Port)
}

func (s *ServerConfig) DisplayAddress() string {
spec, err := s.ListenSpec()
if err != nil {
return s.Address()
}
return spec.DisplayAddress()
}

// DatabaseConfig 数据库连接配置
// 性能优化:新增连接池参数,避免频繁创建/销毁连接
type DatabaseConfig struct {
Expand Down Expand Up @@ -1041,10 +1192,76 @@ type RedisConfig struct {
EnableTLS bool `mapstructure:"enable_tls"`
}

type RedisConnectionNetwork string

const (
RedisConnectionNetworkTCP RedisConnectionNetwork = "tcp"
RedisConnectionNetworkUnix RedisConnectionNetwork = "unix"
)

type RedisConnectionSpec struct {
Network RedisConnectionNetwork
Address string
}

func ParseRedisConnectionSpec(host string, port int) (RedisConnectionSpec, error) {
rawHost := strings.TrimSpace(host)
if rawHost == "" {
rawHost = "localhost"
}

switch {
case strings.HasPrefix(rawHost, "unix:"):
return parseRedisUnixSocketSpec(strings.TrimSpace(strings.TrimPrefix(rawHost, "unix:")))
case strings.HasPrefix(rawHost, "/"):
return parseRedisUnixSocketSpec(rawHost)
default:
if port <= 0 || port > 65535 {
return RedisConnectionSpec{}, fmt.Errorf("tcp redis port must be between 1-65535")
}
return RedisConnectionSpec{
Network: RedisConnectionNetworkTCP,
Address: fmt.Sprintf("%s:%d", rawHost, port),
}, nil
}
}

func parseRedisUnixSocketSpec(raw string) (RedisConnectionSpec, error) {
socketPath := strings.TrimSpace(raw)
if socketPath == "" {
return RedisConnectionSpec{}, fmt.Errorf("redis unix socket path is required")
}
if !strings.HasPrefix(socketPath, "/") {
return RedisConnectionSpec{}, fmt.Errorf("redis unix socket path must be absolute")
}
cleanPath := pathpkg.Clean(socketPath)
if cleanPath == "/" {
return RedisConnectionSpec{}, fmt.Errorf("redis unix socket path cannot be root directory")
}
return RedisConnectionSpec{
Network: RedisConnectionNetworkUnix,
Address: cleanPath,
}, nil
}

func (r *RedisConfig) Address() string {
return fmt.Sprintf("%s:%d", r.Host, r.Port)
}

func (r *RedisConfig) ConnectionSpec() (RedisConnectionSpec, error) {
if r == nil {
return ParseRedisConnectionSpec("", 6379)
}
spec, err := ParseRedisConnectionSpec(r.Host, r.Port)
if err != nil {
return RedisConnectionSpec{}, err
}
if spec.Network == RedisConnectionNetworkUnix && r.EnableTLS {
return RedisConnectionSpec{}, fmt.Errorf("redis.enable_tls is not supported with unix socket connections")
}
return spec, nil
}

type OpsConfig struct {
// Enabled controls whether ops features should run.
//
Expand Down Expand Up @@ -1949,6 +2166,9 @@ func (c *Config) Validate() error {
}
warnIfInsecureURL("server.frontend_url", c.Server.FrontendURL)
}
if _, err := c.Server.ListenSpec(); err != nil {
return fmt.Errorf("server listen config invalid: %w", err)
}
if c.JWT.ExpireHour <= 0 {
return fmt.Errorf("jwt.expire_hour must be positive")
}
Expand Down Expand Up @@ -2190,6 +2410,9 @@ func (c *Config) Validate() error {
if c.Redis.MinIdleConns > c.Redis.PoolSize {
return fmt.Errorf("redis.min_idle_conns cannot exceed redis.pool_size")
}
if _, err := c.Redis.ConnectionSpec(); err != nil {
return fmt.Errorf("redis connection config invalid: %w", err)
}
if c.Dashboard.Enabled {
if c.Dashboard.StatsFreshTTLSeconds <= 0 {
return fmt.Errorf("dashboard_cache.stats_fresh_ttl_seconds must be positive")
Expand Down
Loading
Loading