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
35 changes: 16 additions & 19 deletions internal/config/validation_gateway.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,26 +15,22 @@ func validateGatewayConfig(gateway *StdinGatewayConfig) error {
logValidation.Print("Validating gateway configuration")

// Validate port range using centralized rules
if gateway.Port != nil {
logValidation.Printf("Validating gateway port: %d", *gateway.Port)
if err := PortRange(*gateway.Port, "gateway.port"); err != nil {
return err
}
if err := validateOptionalInt(gateway.Port, "Validating gateway port",
func(v int) *ValidationError { return PortRange(v, "gateway.port") }); err != nil {
return err
}

// Validate timeout values using centralized rules
if gateway.StartupTimeout != nil {
logValidation.Printf("Validating startup timeout: %d", *gateway.StartupTimeout)
if err := TimeoutPositive(*gateway.StartupTimeout, "startupTimeout", "gateway.startupTimeout"); err != nil {
return err
}
if err := validateOptionalInt(gateway.StartupTimeout, "Validating startup timeout",
func(v int) *ValidationError { return TimeoutPositive(v, "startupTimeout", "gateway.startupTimeout") }); err != nil {
return err
}

if gateway.ToolTimeout != nil {
logValidation.Printf("Validating tool timeout: %d", *gateway.ToolTimeout)
if err := TimeoutMinimum(*gateway.ToolTimeout, ToolTimeoutMin, "toolTimeout", "gateway.toolTimeout"); err != nil {
return err
}
if err := validateOptionalInt(gateway.ToolTimeout, "Validating tool timeout",
func(v int) *ValidationError {
return TimeoutMinimum(v, ToolTimeoutMin, "toolTimeout", "gateway.toolTimeout")
}); err != nil {
return err
}

if err := validateContainerRuntimeValue(gateway.ContainerRuntime, "gateway.containerRuntime"); err != nil {
Expand All @@ -57,10 +53,11 @@ func validateGatewayConfig(gateway *StdinGatewayConfig) error {
}

// Validate payloadSizeThreshold per spec §4.1.3.3: must be a positive integer when present.
if gateway.PayloadSizeThreshold != nil {
if err := validateGatewayPayloadSizeThreshold(*gateway.PayloadSizeThreshold, "payloadSizeThreshold", "gateway.payloadSizeThreshold"); err != nil {
return err
}
if err := validateOptionalInt(gateway.PayloadSizeThreshold, "Validating payload size threshold",
func(v int) *ValidationError {
return PositiveInteger(v, "payloadSizeThreshold", "gateway.payloadSizeThreshold")
}); err != nil {
return err
}

// Validate trustedBots per spec §4.1.3.4: must be non-empty array when present
Expand Down
27 changes: 27 additions & 0 deletions internal/config/validation_rules.go
Original file line number Diff line number Diff line change
Expand Up @@ -205,6 +205,33 @@ func NonEmptyString(value, fieldName, jsonPath string) *ValidationError {
fmt.Sprintf("Provide a non-empty value for %s", fieldName))
}

// validateOptionalInt nil-guards an optional integer pointer field before
// calling validateFn with the dereferenced value. It replaces the repeated
// boilerplate:
//
// if ptr != nil {
// logValidation.Print(logMsg)
// if err := someRule(*ptr, ...); err != nil { return err }
// }
//
// validateFn receives the dereferenced value and returns a *ValidationError
// (nil means valid). The caller controls all validation logic through the
// closure it passes. logMsg is emitted at debug level only when ptr is
// non-nil. Note: the rule functions called inside validateFn already log the
// field name, value, and JSON path, so logMsg is intentionally a short
// higher-level label ("Validating gateway port") rather than a formatted string
// containing the value.
func validateOptionalInt(ptr *int, logMsg string, validateFn func(int) *ValidationError) error {
if ptr == nil {
return nil
}
logValidation.Print(logMsg)
if ve := validateFn(*ptr); ve != nil {
return ve
}
return nil
}

// AbsolutePath validates that a directory path is an absolute path
// Per MCP Gateway schema: Unix paths start with '/', Windows paths start with a drive letter followed by ':\'
// Pattern: ^(/|[A-Za-z]:\\)
Expand Down
62 changes: 62 additions & 0 deletions internal/config/validation_rules_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1414,3 +1414,65 @@ func TestTimeoutRange(t *testing.T) {
})
}
}

func TestValidateOptionalInt(t *testing.T) {
tests := []struct {
name string
ptr *int
validateF func(int) *ValidationError
wantErr bool
errMsg string
}{
{
name: "nil pointer returns nil without calling validateFn",
ptr: nil,
validateF: func(int) *ValidationError { panic("should not be called") },
wantErr: false,
},
{
name: "non-nil valid value returns nil",
ptr: intPtr(80),
validateF: func(v int) *ValidationError { return PortRange(v, "gateway.port") },
wantErr: false,
},
{
name: "non-nil invalid value returns error",
ptr: intPtr(0),
validateF: func(v int) *ValidationError { return PortRange(v, "gateway.port") },
wantErr: true,
errMsg: "port must be between 1 and 65535",
},
{
name: "closure can embed additional guard (zero means skip)",
ptr: intPtr(0),
validateF: func(v int) *ValidationError {
if v == 0 {
return nil // sentinel: skip validation
}
return TimeoutMinimum(v, 10, "toolTimeout", "mcpServers.s.toolTimeout")
},
wantErr: false,
},
{
name: "closure returns error when value is below minimum",
ptr: intPtr(5),
validateF: func(v int) *ValidationError {
return TimeoutMinimum(v, 10, "toolTimeout", "mcpServers.s.toolTimeout")
},
wantErr: true,
errMsg: "toolTimeout must be at least 10",
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := validateOptionalInt(tt.ptr, "test log message", tt.validateF)
if tt.wantErr {
require.Error(t, err)
assert.Contains(t, err.Error(), tt.errMsg)
} else {
assert.NoError(t, err)
}
})
}
}
18 changes: 12 additions & 6 deletions internal/config/validation_server.go
Original file line number Diff line number Diff line change
Expand Up @@ -85,12 +85,18 @@ func validateStandardServerConfig(name string, server *StdinServerConfig, jsonPa

// Validate per-server toolTimeout if provided and non-zero.
// A value of 0 means "unset – fall back to the global gateway timeout".
if server.ToolTimeout != nil && *server.ToolTimeout != 0 {
toolTimeoutField := server.toolTimeoutField()
if err := TimeoutMinimum(*server.ToolTimeout, ToolTimeoutMin, toolTimeoutField, jsonPath+"."+toolTimeoutField); err != nil {
return logValidationFail(
name, server.Type, fmt.Sprintf("%s %d is below minimum %d", toolTimeoutField, *server.ToolTimeout, ToolTimeoutMin), err)
}
toolTimeoutField := server.toolTimeoutField()
if err := validateOptionalInt(server.ToolTimeout, "Validating per-server tool timeout",
func(v int) *ValidationError {
if v == 0 {
return nil // 0 means inherit from global gateway timeout
}
return TimeoutMinimum(v, ToolTimeoutMin, toolTimeoutField, jsonPath+"."+toolTimeoutField)
}); err != nil {
// server.ToolTimeout is non-nil here: validateOptionalInt only invokes the
// callback (and can only return an error) when ptr is non-nil.
return logValidationFail(
name, server.Type, fmt.Sprintf("%s %d is below minimum %d", toolTimeoutField, *server.ToolTimeout, ToolTimeoutMin), err)
}

if err := validateCommonServerFields(name, server.Type, server.Auth, server.ToolResponseFilters, jsonPath); err != nil {
Expand Down
Loading