Explorar el Código

validate toc & api listener config

Mike hace 11 meses
padre
commit
1c8d863180
Se han modificado 5 ficheros con 266 adiciones y 38 borrados
  1. 14 10
      cmd/server/factory.go
  2. 1 1
      cmd/server/main.go
  3. 66 17
      config/config.go
  4. 177 1
      config/config_test.go
  5. 8 9
      server/kerberos/kerberos.go

+ 14 - 10
cmd/server/factory.go

@@ -41,27 +41,31 @@ func MakeCommonDeps() (Container, error) {
 	c := Container{}
 
 	if err := validateConfigMigration(); err != nil {
-		return c, fmt.Errorf("unable to validate config migration: %s\n", err.Error())
+		return c, fmt.Errorf("unable to validate config migration: %s", err.Error())
 	}
 
 	err := envconfig.Process("", &c.cfg)
 	if err != nil {
-		return c, fmt.Errorf("unable to process app config: %s\n", err.Error())
+		return c, fmt.Errorf("unable to process app config: %s", err.Error())
 	}
 
-	c.Listeners, err = config.ParseListenersCfg(c.cfg.BOSListeners, c.cfg.BOSAdvertisedHosts, c.cfg.KerberosListeners)
+	if err := c.cfg.Validate(); err != nil {
+		return c, fmt.Errorf("configuration validation failed: %s", err.Error())
+	}
+
+	c.Listeners, err = c.cfg.ParseListenersCfg()
 	if err != nil {
-		return c, fmt.Errorf("unable to parse listener config: %s\n", err.Error())
+		return c, fmt.Errorf("unable to parse listener config: %s", err.Error())
 	}
 
 	c.sqLiteUserStore, err = state.NewSQLiteUserStore(c.cfg.DBPath)
 	if err != nil {
-		return c, fmt.Errorf("unable to create feedbag store: %s\n", err.Error())
+		return c, fmt.Errorf("unable to create feedbag store: %s", err.Error())
 	}
 
 	c.hmacCookieBaker, err = state.NewHMACCookieBaker()
 	if err != nil {
-		return c, fmt.Errorf("unable to create HMAC cookie baker: %s\n", err.Error())
+		return c, fmt.Errorf("unable to create HMAC cookie baker: %s", err.Error())
 	}
 
 	c.logger = middleware.NewLogger(c.cfg)
@@ -146,17 +150,17 @@ func validateConfigMigration() error {
 			if contains(newEnvVarsMissing, "OSCAR_ADVERTISED_LISTENERS") {
 				oscarHost := getEnvOrDefault("OSCAR_HOST", "127.0.0.1")
 				authPort := getEnvOrDefault("AUTH_PORT", "5190")
-				errorMsg.WriteString(fmt.Sprintf("export OSCAR_ADVERTISED_LISTENERS=EXTERNAL://%s:%s\n", oscarHost, authPort))
+				errorMsg.WriteString(fmt.Sprintf("export OSCAR_ADVERTISED_LISTENERS=LOCAL://%s:%s\n", oscarHost, authPort))
 			}
 
 			if contains(newEnvVarsMissing, "OSCAR_LISTENERS") {
 				authPort := getEnvOrDefault("AUTH_PORT", "5190")
-				errorMsg.WriteString(fmt.Sprintf("export OSCAR_LISTENERS=EXTERNAL://0.0.0.0:%s\n", authPort))
+				errorMsg.WriteString(fmt.Sprintf("export OSCAR_LISTENERS=LOCAL://0.0.0.0:%s\n", authPort))
 			}
 
 			if contains(newEnvVarsMissing, "KERBEROS_LISTENERS") {
 				kerberosPort := getEnvOrDefault("KERBEROS_PORT", "1088")
-				errorMsg.WriteString(fmt.Sprintf("export KERBEROS_LISTENERS=EXTERNAL://0.0.0.0:%s\n", kerberosPort))
+				errorMsg.WriteString(fmt.Sprintf("export KERBEROS_LISTENERS=LOCAL://0.0.0.0:%s\n", kerberosPort))
 			}
 
 			if contains(newEnvVarsMissing, "TOC_LISTENERS") {
@@ -317,7 +321,7 @@ func OSCAR(deps Container) *oscar.Server {
 
 // KerberosAPI creates an HTTP server for the Kerberos server.
 func KerberosAPI(deps Container) *kerberos.Server {
-	logger := deps.logger.With("svc", "kerberos")
+	logger := deps.logger.With("svc", "Kerberos")
 	authService := foodgroup.NewAuthService(deps.cfg, deps.inMemorySessionManager, deps.inMemorySessionManager, deps.chatSessionManager, deps.sqLiteUserStore, deps.hmacCookieBaker, deps.chatSessionManager, deps.sqLiteUserStore, deps.rateLimitClasses)
 	return kerberos.NewKerberosServer(deps.Listeners, logger, authService)
 }

+ 1 - 1
cmd/server/main.go

@@ -52,7 +52,7 @@ func main() {
 
 	deps, err := MakeCommonDeps()
 	if err != nil {
-		fmt.Printf("error initializing common deps: %v\n", err)
+		fmt.Printf("startup failed: %s\n", err)
 		os.Exit(1)
 	}
 

+ 66 - 17
config/config.go

@@ -25,6 +25,18 @@ func (e uriFormatError) Error() string {
 	return fmt.Sprintf("invalid listener URI %q: %v. Valid format: SCHEME://HOST:PORT (e.g., LOCAL://0.0.0.0:5190)", e.URI, e.Err)
 }
 
+type Build struct {
+	Version string `json:"version"`
+	Commit  string `json:"commit"`
+	Date    string `json:"date"`
+}
+
+type Listener struct {
+	BOSListenAddress      string
+	BOSAdvertisedHost     string
+	KerberosListenAddress string
+}
+
 //go:generate go run ../cmd/config_generator unix settings.env ssl
 type Config struct {
 	BOSListeners       string `envconfig:"OSCAR_LISTENERS" required:"true" basic:"LOCAL://0.0.0.0:5190" ssl:"PLAINTEXT://0.0.0.0:5190,SSL://0.0.0.0:5192" description:"Network listeners for core OSCAR services. For multi-homed servers, allows users to connect from multiple networks. For example, you can allow both LAN and Internet clients to connect to the same server using different connection settings.\n\nFormat:\n\t- Comma-separated list of [NAME]://[HOSTNAME]:[PORT]\n\t- Listener names and ports must be unique\n\t- Listener names are user-defined\n\t- Each listener needs OSCAR_ADVERTISED_LISTENERS/KERBEROS_LISTENERS configs\n\nExamples:\n\t// Listen on all interfaces\n\tLAN://0.0.0.0:5190\n\t// Separate Internet and LAN config\n\tWAN://142.250.176.206:5190,LAN://192.168.1.10:5191"`
@@ -38,19 +50,7 @@ type Config struct {
 	LogLevel    string `envconfig:"LOG_LEVEL" required:"true" basic:"info" ssl:"info" description:"Set logging granularity. Possible values: 'trace', 'debug', 'info', 'warn', 'error'."`
 }
 
-type Build struct {
-	Version string `json:"version"`
-	Commit  string `json:"commit"`
-	Date    string `json:"date"`
-}
-
-type Listener struct {
-	BOSListenAddress      string
-	BOSAdvertisedHost     string
-	KerberosListenAddress string
-}
-
-func ParseListenersCfg(BOSListeners string, BOSAdvertisedListeners string, kerberosListeners string) ([]Listener, error) {
+func (c *Config) ParseListenersCfg() ([]Listener, error) {
 	// Helper function to parse and validate a single URI
 	parseURI := func(uriStr string) (*url.URL, error) {
 		uriStr = strings.TrimSpace(uriStr)
@@ -62,8 +62,13 @@ func ParseListenersCfg(BOSListeners string, BOSAdvertisedListeners string, kerbe
 		if err != nil {
 			return nil, uriFormatError{URI: uriStr, Err: err}
 		}
-		if u.Scheme == "" {
+		switch {
+		case u.Scheme == "":
 			return nil, uriFormatError{URI: uriStr, Err: errors.New("missing scheme")}
+		case u.Hostname() == "":
+			return nil, uriFormatError{URI: uriStr, Err: errors.New("missing host")}
+		case u.Port() == "":
+			return nil, uriFormatError{URI: uriStr, Err: errors.New("missing port")}
 		}
 
 		return u, nil
@@ -72,7 +77,7 @@ func ParseListenersCfg(BOSListeners string, BOSAdvertisedListeners string, kerbe
 	m := make(map[string]*Listener)
 
 	// Parse BOS listeners
-	for _, uriStr := range strings.Split(BOSListeners, ",") {
+	for _, uriStr := range strings.Split(c.BOSListeners, ",") {
 		u, err := parseURI(uriStr)
 		if err != nil {
 			return nil, err
@@ -91,7 +96,7 @@ func ParseListenersCfg(BOSListeners string, BOSAdvertisedListeners string, kerbe
 	}
 
 	// Parse BOS advertised listeners
-	for _, uriStr := range strings.Split(BOSAdvertisedListeners, ",") {
+	for _, uriStr := range strings.Split(c.BOSAdvertisedHosts, ",") {
 		u, err := parseURI(uriStr)
 		if err != nil {
 			return nil, err
@@ -110,7 +115,7 @@ func ParseListenersCfg(BOSListeners string, BOSAdvertisedListeners string, kerbe
 	}
 
 	// Parse Kerberos listeners
-	for _, uriStr := range strings.Split(kerberosListeners, ",") {
+	for _, uriStr := range strings.Split(c.KerberosListeners, ",") {
 		u, err := parseURI(uriStr)
 		if err != nil {
 			return nil, err
@@ -146,3 +151,47 @@ func ParseListenersCfg(BOSListeners string, BOSAdvertisedListeners string, kerbe
 
 	return ret, nil
 }
+
+func (c *Config) Validate() error {
+	// Validate TOCListeners (format: hostname:port pairs)
+	for _, listener := range strings.Split(c.TOCListeners, ",") {
+		listener = strings.TrimSpace(listener)
+		if listener == "" {
+			continue
+		}
+
+		host, port, err := net.SplitHostPort(listener)
+		if err != nil {
+			return fmt.Errorf("invalid TOC listener %q: %v. Valid format: HOST:PORT (e.g., 0.0.0.0:9898)", listener, err)
+		}
+
+		if host == "" {
+			return fmt.Errorf("invalid TOC listener %q: missing host. Valid format: HOST:PORT (e.g., 0.0.0.0:9898)", listener)
+		}
+
+		if port == "" {
+			return fmt.Errorf("invalid TOC listener %q: missing port. Valid format: HOST:PORT (e.g., 0.0.0.0:9898)", listener)
+		}
+	}
+
+	// Validate APIListener (format: hostname:port pair, no scheme)
+	apiListener := strings.TrimSpace(c.APIListener)
+	if apiListener == "" {
+		return fmt.Errorf("APIListener is required and cannot be empty")
+	}
+
+	host, port, err := net.SplitHostPort(apiListener)
+	if err != nil {
+		return fmt.Errorf("invalid API listener %q: %v. Valid format: HOST:PORT (e.g., 127.0.0.1:8080)", c.APIListener, err)
+	}
+
+	if host == "" {
+		return fmt.Errorf("invalid API listener %q: missing host. Valid format: HOST:PORT (e.g., 127.0.0.1:8080)", c.APIListener)
+	}
+
+	if port == "" {
+		return fmt.Errorf("invalid API listener %q: missing port. Valid format: HOST:PORT (e.g., 127.0.0.1:8080)", c.APIListener)
+	}
+
+	return nil
+}

+ 177 - 1
config/config_test.go

@@ -142,6 +142,24 @@ func TestParseListenersCfg(t *testing.T) {
 			wantErr:                true,
 			errContains:            "Valid format: SCHEME://HOST:PORT",
 		},
+		{
+			name:                   "BOS listener missing port",
+			bosListeners:           "LOCAL://0.0.0.0",
+			bosAdvertisedListeners: "LOCAL://127.0.0.1:5190",
+			kerberosListeners:      "",
+			want:                   nil,
+			wantErr:                true,
+			errContains:            "missing port",
+		},
+		{
+			name:                   "BOS listener missing host",
+			bosListeners:           "LOCAL://:5190",
+			bosAdvertisedListeners: "LOCAL://127.0.0.1:5190",
+			kerberosListeners:      "",
+			want:                   nil,
+			wantErr:                true,
+			errContains:            "missing host",
+		},
 		{
 			name:                   "complex multi-listener setup",
 			bosListeners:           "LAN://192.168.1.10:5190,WAN://0.0.0.0:5191,DOCKER://172.17.0.1:5192",
@@ -197,7 +215,12 @@ func TestParseListenersCfg(t *testing.T) {
 
 	for _, tt := range tests {
 		t.Run(tt.name, func(t *testing.T) {
-			got, err := ParseListenersCfg(tt.bosListeners, tt.bosAdvertisedListeners, tt.kerberosListeners)
+			config := &Config{
+				BOSListeners:       tt.bosListeners,
+				BOSAdvertisedHosts: tt.bosAdvertisedListeners,
+				KerberosListeners:  tt.kerberosListeners,
+			}
+			got, err := config.ParseListenersCfg()
 
 			if tt.wantErr {
 				if err == nil {
@@ -255,6 +278,159 @@ func TestParseListenersCfg(t *testing.T) {
 	}
 }
 
+func TestConfigValidate(t *testing.T) {
+	tests := []struct {
+		name        string
+		config      Config
+		wantErr     bool
+		errContains string
+	}{
+		{
+			name: "valid config with all fields",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898,192.168.1.10:9899",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr: false,
+		},
+		{
+			name: "valid config with single TOC listener",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr: false,
+		},
+		{
+			name: "valid config with empty TOC listeners",
+			config: Config{
+				TOCListeners: "",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr: false,
+		},
+		{
+			name: "valid config with empty API listener",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898",
+				APIListener:  "",
+			},
+			wantErr:     true,
+			errContains: "APIListener is required and cannot be empty",
+		},
+		{
+			name: "valid config with all empty",
+			config: Config{
+				TOCListeners: "",
+				APIListener:  "",
+			},
+			wantErr:     true,
+			errContains: "APIListener is required and cannot be empty",
+		},
+		{
+			name: "invalid TOC listener - missing port",
+			config: Config{
+				TOCListeners: "0.0.0.0",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr:     true,
+			errContains: "invalid TOC listener \"0.0.0.0\": address 0.0.0.0: missing port in address",
+		},
+		{
+			name: "invalid TOC listener - missing host",
+			config: Config{
+				TOCListeners: ":9898",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr:     true,
+			errContains: "invalid TOC listener \":9898\": missing host",
+		},
+		{
+			name: "invalid TOC listener - malformed",
+			config: Config{
+				TOCListeners: "invalid-format",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr:     true,
+			errContains: "invalid TOC listener \"invalid-format\": address invalid-format: missing port in address",
+		},
+		{
+			name: "invalid TOC listener in comma-separated list",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898,invalid-format,192.168.1.10:9899",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr:     true,
+			errContains: "invalid TOC listener \"invalid-format\": address invalid-format: missing port in address",
+		},
+		{
+			name: "invalid API listener - missing port",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898",
+				APIListener:  "127.0.0.1",
+			},
+			wantErr:     true,
+			errContains: "invalid API listener \"127.0.0.1\": address 127.0.0.1: missing port in address",
+		},
+		{
+			name: "invalid API listener - missing host",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898",
+				APIListener:  ":8080",
+			},
+			wantErr:     true,
+			errContains: "invalid API listener \":8080\": missing host",
+		},
+		{
+			name: "invalid API listener - malformed",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898",
+				APIListener:  "invalid-format",
+			},
+			wantErr:     true,
+			errContains: "invalid API listener \"invalid-format\": address invalid-format: missing port in address",
+		},
+		{
+			name: "whitespace-only TOC listeners",
+			config: Config{
+				TOCListeners: "   ,  ,  ",
+				APIListener:  "127.0.0.1:8080",
+			},
+			wantErr: false,
+		},
+		{
+			name: "whitespace-only API listener",
+			config: Config{
+				TOCListeners: "0.0.0.0:9898",
+				APIListener:  "   ",
+			},
+			wantErr:     true,
+			errContains: "APIListener is required and cannot be empty",
+		},
+	}
+
+	for _, tt := range tests {
+		t.Run(tt.name, func(t *testing.T) {
+			err := tt.config.Validate()
+
+			if tt.wantErr {
+				if err == nil {
+					t.Errorf("Config.Validate() expected error but got none")
+					return
+				}
+				if tt.errContains != "" && !contains(err.Error(), tt.errContains) {
+					t.Errorf("Config.Validate() error = %v, want error containing %q", err, tt.errContains)
+				}
+				return
+			}
+
+			if err != nil {
+				t.Errorf("Config.Validate() unexpected error = %v", err)
+			}
+		})
+	}
+}
+
 // Helper function to check if a string contains a substring
 func contains(s, substr string) bool {
 	return len(s) >= len(substr) && (s == substr || len(substr) == 0 ||

+ 8 - 9
server/kerberos/kerberos.go

@@ -7,7 +7,6 @@ import (
 	"fmt"
 	"io"
 	"log/slog"
-	"net"
 	"net/http"
 
 	"golang.org/x/sync/errgroup"
@@ -50,14 +49,13 @@ func NewKerberosServer(listeners []config.Listener, logger *slog.Logger, authSer
 // Server hosts an HTTP endpoint capable of handling AIM-style Kerberos
 // authentication. The messages are structured as SNACs transmitted over HTTP.
 type Server struct {
-	servers   []*http.Server
-	listeners []net.Listener
-	logger    *slog.Logger
+	servers []*http.Server
+	logger  *slog.Logger
 }
 
 func (s *Server) ListenAndServe() error {
 	if len(s.servers) == 0 {
-		s.logger.Info("no kerberos listeners defined, moving on")
+		s.logger.Debug("no kerberos listeners defined")
 		return nil
 	}
 
@@ -80,10 +78,11 @@ func (s *Server) ListenAndServe() error {
 }
 
 func (s *Server) Shutdown(ctx context.Context) error {
-	defer s.logger.Info("shutdown complete")
-
-	for _, srv := range s.servers {
-		_ = srv.Shutdown(ctx)
+	if len(s.servers) > 0 {
+		for _, srv := range s.servers {
+			_ = srv.Shutdown(ctx)
+		}
+		s.logger.Info("shutdown complete")
 	}
 	return nil
 }