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
5 changes: 4 additions & 1 deletion config.yaml.example
Original file line number Diff line number Diff line change
Expand Up @@ -149,9 +149,12 @@ websocket:
# Within this rate budget, a full concurrent cap accepts then closes with 1013
# before hello, allowing browsers to back off without evicting existing clients.
max_connects_per_minute: 10
# Other sites allowed to open /ws. Exact scheme://host[:port], no wildcards.
# Other sites allowed to open /ws. Exact scheme://host[:port], or
# scheme://*.domain[:port] for every subdomain of domain at any depth. The
# wildcard never matches the apex itself, and scheme and port stay exact.
#allowed_origins:
# - https://example.com
# - https://*.example.com

# Node staleness, deletion, and clock-drift thresholds.
#nodes:
Expand Down
18 changes: 14 additions & 4 deletions internal/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"net/url"
"os"
"path/filepath"
"slices"
"strings"
"time"

Expand Down Expand Up @@ -243,7 +244,8 @@ type WebSocketConfig struct {
// MaxConnectsPerMinute limits upgrade attempts, including failed handshakes.
// Zero/omitted defaults to 10; IPv6 addresses share a /64 attempt budget.
MaxConnectsPerMinute int `yaml:"max_connects_per_minute"`
// AllowedOrigins are extra exact origins allowed to open /ws. Defaults to same-host only.
// AllowedOrigins are extra origins allowed to open /ws: exact, or scheme://*.domain[:port]
// for every subdomain of domain (never the apex). Defaults to same-host only.
AllowedOrigins []string `yaml:"allowed_origins"`
}

Expand Down Expand Up @@ -415,16 +417,24 @@ func Load(path string) (*Config, error) {
return cfg, nil
}

// validateOrigin rejects wildcards, which the WebSocket library would treat as patterns.
// validateOrigin allows exact origins or a leading "*." subdomain wildcard; the WebSocket
// library treats entries as path.Match patterns, so every other pattern character is rejected.
func validateOrigin(origin string) error {
u, err := url.Parse(origin)
if err != nil {
return fmt.Errorf("invalid origin %q: %w", origin, err)
}
host := u.Hostname()
// At least two labels after the wildcard, so it never spans a whole TLD.
if rest, ok := strings.CutPrefix(host, "*."); ok {
if labels := strings.Split(rest, "."); len(labels) >= 2 && !slices.Contains(labels, "") {
host = rest
}
}
if (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.User != nil ||
u.Path != "" || u.RawQuery != "" || u.Fragment != "" || u.Opaque != "" ||
strings.ContainsAny(origin, `*?[]\`) {
return fmt.Errorf("origin %q must be exactly scheme://host[:port] with an http or https scheme", origin)
strings.Contains(host, "*") || strings.ContainsAny(origin, `?[]\`) {
return fmt.Errorf("origin %q must be scheme://host[:port] or scheme://*.domain[:port] with an http or https scheme", origin)
}
return nil
}
Expand Down
13 changes: 12 additions & 1 deletion internal/config/ws_connect_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,18 @@ func TestWebSocketAllowedOrigins(t *testing.T) {
{"https", "https://example.com", false},
{"http with port", "http://localhost:5173", false},
{"wildcard", "*", true},
{"wildcard host", "https://*.example.com", true},
{"wildcard subdomain", "https://*.example.com", false},
{"wildcard subdomain with port", "http://*.example.com:8443", false},
{"wildcard scheme only", "https://*", true},
{"wildcard over tld", "https://*.com", true},
{"wildcard label prefix", "https://a*.example.com", true},
{"wildcard label suffix", "https://*a.example.com", true},
{"wildcard inner label", "https://x.*.example.com", true},
{"double wildcard", "https://*.*.example.com", true},
{"wildcard empty label", "https://*..com", true},
{"question mark", "https://ex?mple.com", true},
{"character class", "https://[ab].example.com", true},
{"backslash", `https://*.example.com\`, true},
{"no scheme", "example.com", true},
{"other scheme", "ftp://example.com", true},
{"path", "https://example.com/", true},
Expand Down
12 changes: 12 additions & 0 deletions internal/ws/handler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@ import (
)

func TestAllowedOrigins(t *testing.T) {
wildcard := []string{"https://*.example.com"}
apexAndWildcard := []string{"https://example.com", "https://*.example.com"}
for _, tc := range []struct {
name string
allowed []string
Expand All @@ -30,6 +32,16 @@ func TestAllowedOrigins(t *testing.T) {
{"scheme mismatch", []string{"https://example.com"}, "http://example.com", false},
{"port mismatch", []string{"https://example.com"}, "https://example.com:8443", false},
{"unlisted origin", []string{"https://example.com"}, "https://other.example", false},
{"wildcard subdomain", wildcard, "https://sub.example.com", true},
{"wildcard deeper subdomain", wildcard, "https://a.b.example.com", true},
{"wildcard other case", wildcard, "https://SUB.Example.com", true},
{"wildcard apex", wildcard, "https://example.com", false},
{"wildcard lookalike", wildcard, "https://evilexample.com", false},
{"wildcard suffix host", wildcard, "https://example.com.evil.net", false},
{"wildcard scheme mismatch", wildcard, "http://sub.example.com", false},
{"wildcard port mismatch", wildcard, "https://sub.example.com:8443", false},
{"apex and wildcard, apex", apexAndWildcard, "https://example.com", true},
{"apex and wildcard, subdomain", apexAndWildcard, "https://sub.example.com", true},
} {
t.Run(tc.name, func(t *testing.T) {
server := httptest.NewServer(Handler(hub.New(), nil, 5, 100, tc.allowed))
Expand Down
Loading