Merge pull request #7 from wahyd4/wildcard-domain

Support wildcard domain
This commit is contained in:
2025-08-04 23:27:35 +10:00
committed by GitHub
10 changed files with 426 additions and 195 deletions
+2
View File
@@ -70,6 +70,8 @@ webauthn:
cors:
allowed_origins:
- "https://your-domain.com" # Your domain with protocol
# Wildcard domains are supported for subdomains:
- "*.your-domain.com" # Matches api.your-domain.com, app.your-domain.com, etc.
auth:
session_secret: "your-secure-secret-key" # Generate a secure random string
+89
View File
@@ -0,0 +1,89 @@
package main
import (
"net/http"
"net/http/httptest"
"testing"
"passkey-auth/internal/cors"
)
func TestWildcardCORSIntegration(t *testing.T) {
// Create a simple test handler
testHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
w.Write([]byte("OK"))
})
// Create CORS middleware with wildcard support
corsMiddleware := cors.WildcardCORS(cors.Config{
AllowedOrigins: []string{"*.junv.cc", "https://static.example.com"},
AllowedMethods: []string{"GET", "POST", "OPTIONS"},
AllowedHeaders: []string{"*"},
AllowCredentials: true,
})
// Wrap the test handler
handler := corsMiddleware(testHandler)
tests := []struct {
name string
origin string
expectAllowed bool
expectedOrigin string
}{
{
name: "wildcard subdomain match",
origin: "https://api.junv.cc",
expectAllowed: true,
expectedOrigin: "https://api.junv.cc",
},
{
name: "wildcard base domain match",
origin: "https://junv.cc",
expectAllowed: true,
expectedOrigin: "https://junv.cc",
},
{
name: "static domain match",
origin: "https://static.example.com",
expectAllowed: true,
},
{
name: "no match",
origin: "https://evil.com",
expectAllowed: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Create a preflight OPTIONS request
req := httptest.NewRequest("OPTIONS", "/", nil)
req.Header.Set("Origin", tt.origin)
req.Header.Set("Access-Control-Request-Method", "POST")
// Record the response
w := httptest.NewRecorder()
handler.ServeHTTP(w, req)
// Check CORS headers
allowOriginHeader := w.Header().Get("Access-Control-Allow-Origin")
if tt.expectAllowed {
if allowOriginHeader == "" {
t.Errorf("Expected Access-Control-Allow-Origin header, but got none")
}
// For wildcard matches, should return the specific origin
if tt.expectedOrigin != "" && allowOriginHeader != tt.expectedOrigin {
t.Errorf("Expected Access-Control-Allow-Origin: %s, got: %s", tt.expectedOrigin, allowOriginHeader)
}
} else {
if allowOriginHeader != "" {
t.Errorf("Expected no Access-Control-Allow-Origin header, but got: %s", allowOriginHeader)
}
}
})
}
}
+68
View File
@@ -0,0 +1,68 @@
package cors
import (
"net/http"
"strings"
"github.com/rs/cors"
)
// Config holds the CORS configuration with wildcard support
type Config struct {
AllowedOrigins []string
AllowedMethods []string
AllowedHeaders []string
AllowCredentials bool
}
// WildcardCORS creates a CORS handler with wildcard domain support
func WildcardCORS(config Config) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
origin := r.Header.Get("Origin")
// Determine allowed origins for this request
var allowedOrigins []string
// Separate wildcard and static origins
var wildcardPatterns []string
var staticOrigins []string
for _, configuredOrigin := range config.AllowedOrigins {
if strings.Contains(configuredOrigin, "*") {
wildcardPatterns = append(wildcardPatterns, configuredOrigin)
} else {
staticOrigins = append(staticOrigins, configuredOrigin)
}
}
// Check if origin matches any wildcard pattern
if len(wildcardPatterns) > 0 && origin != "" {
wildcardMatcher := NewWildcardMatcher(wildcardPatterns)
if wildcardMatcher.MatchOrigin(origin) {
// For wildcard matches, allow the specific origin
allowedOrigins = []string{origin}
}
}
// If no wildcard match, use static origins
if len(allowedOrigins) == 0 {
allowedOrigins = staticOrigins
} else {
// If we had a wildcard match, also include static origins
allowedOrigins = append(allowedOrigins, staticOrigins...)
}
// Create a new CORS instance for this request with the determined origins
c := cors.New(cors.Options{
AllowedOrigins: allowedOrigins,
AllowedMethods: config.AllowedMethods,
AllowedHeaders: config.AllowedHeaders,
AllowCredentials: config.AllowCredentials,
})
// Use the rs/cors handler
c.Handler(next).ServeHTTP(w, r)
})
}
}
+100
View File
@@ -0,0 +1,100 @@
package cors
import (
"strings"
)
// WildcardMatcher provides wildcard domain matching for CORS origins
type WildcardMatcher struct {
patterns []string
}
// NewWildcardMatcher creates a new wildcard matcher with the given patterns
func NewWildcardMatcher(patterns []string) *WildcardMatcher {
return &WildcardMatcher{
patterns: patterns,
}
}
// MatchOrigin checks if the given origin matches any of the wildcard patterns
func (m *WildcardMatcher) MatchOrigin(origin string) bool {
for _, pattern := range m.patterns {
if m.matchPattern(origin, pattern) {
return true
}
}
return false
}
// matchPattern checks if origin matches a specific pattern
// Supports patterns like:
// - "*.example.com" matches "api.example.com", "auth.example.com", etc.
// - "*.*.example.com" matches "api.v1.example.com", etc.
// - "example.com" matches exactly "example.com"
func (m *WildcardMatcher) matchPattern(origin, pattern string) bool {
// Remove protocol from origin if present
origin = strings.TrimPrefix(origin, "https://")
origin = strings.TrimPrefix(origin, "http://")
// Remove port if present
if colonIndex := strings.LastIndex(origin, ":"); colonIndex != -1 && colonIndex > strings.LastIndex(origin, "]") {
origin = origin[:colonIndex]
}
// Exact match
if origin == pattern {
return true
}
// Wildcard match
if strings.Contains(pattern, "*") {
return m.wildcardMatch(origin, pattern)
}
return false
}
// wildcardMatch performs wildcard matching
func (m *WildcardMatcher) wildcardMatch(origin, pattern string) bool {
// Handle simple case: *.domain.com
if strings.HasPrefix(pattern, "*.") {
suffix := pattern[2:] // Remove "*."
// Check if origin ends with the suffix and has at least one subdomain
if strings.HasSuffix(origin, "."+suffix) {
// Ensure there's a subdomain (not just the suffix itself)
prefix := strings.TrimSuffix(origin, "."+suffix)
// Make sure the prefix doesn't contain dots (single-level subdomain wildcard)
// If you want multi-level subdomains, remove this check
return !strings.Contains(prefix, ".")
}
// Also check if origin exactly matches the suffix (without subdomain)
return origin == suffix
}
// For more complex patterns, we could implement more sophisticated matching
// For now, handle the common *.domain.com case
return false
}
// GetAllowedOrigins returns the actual allowed origins for a request
// This expands wildcard patterns based on the request origin
func (m *WildcardMatcher) GetAllowedOrigins(requestOrigin string, staticOrigins []string) []string {
allowedOrigins := make([]string, 0, len(staticOrigins))
for _, origin := range staticOrigins {
if strings.Contains(origin, "*") {
// This is a wildcard pattern
if m.matchPattern(requestOrigin, origin) {
// Add the actual request origin instead of the pattern
allowedOrigins = append(allowedOrigins, requestOrigin)
}
} else {
// This is a static origin, add as-is
allowedOrigins = append(allowedOrigins, origin)
}
}
return allowedOrigins
}
+138
View File
@@ -0,0 +1,138 @@
package cors
import (
"testing"
)
func TestWildcardMatcher(t *testing.T) {
tests := []struct {
name string
patterns []string
origin string
expected bool
}{
{
name: "exact match",
patterns: []string{"example.com"},
origin: "example.com",
expected: true,
},
{
name: "exact match with https",
patterns: []string{"example.com"},
origin: "https://example.com",
expected: true,
},
{
name: "wildcard subdomain match",
patterns: []string{"*.junv.cc"},
origin: "api.junv.cc",
expected: true,
},
{
name: "wildcard subdomain match with https",
patterns: []string{"*.junv.cc"},
origin: "https://auth.junv.cc",
expected: true,
},
{
name: "wildcard subdomain match with port",
patterns: []string{"*.junv.cc"},
origin: "https://dev.junv.cc:3000",
expected: true,
},
{
name: "wildcard base domain match",
patterns: []string{"*.junv.cc"},
origin: "junv.cc",
expected: true,
},
{
name: "wildcard no match - different domain",
patterns: []string{"*.junv.cc"},
origin: "api.example.com",
expected: false,
},
{
name: "wildcard no match - multi-level subdomain",
patterns: []string{"*.junv.cc"},
origin: "api.v1.junv.cc",
expected: false,
},
{
name: "multiple patterns - first match",
patterns: []string{"*.junv.cc", "*.example.com"},
origin: "api.junv.cc",
expected: true,
},
{
name: "multiple patterns - second match",
patterns: []string{"*.junv.cc", "*.example.com"},
origin: "api.example.com",
expected: true,
},
{
name: "no match",
patterns: []string{"*.junv.cc", "*.example.com"},
origin: "api.other.com",
expected: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
matcher := NewWildcardMatcher(tt.patterns)
result := matcher.MatchOrigin(tt.origin)
if result != tt.expected {
t.Errorf("MatchOrigin() = %v, expected %v for origin %s with patterns %v",
result, tt.expected, tt.origin, tt.patterns)
}
})
}
}
func TestGetAllowedOrigins(t *testing.T) {
tests := []struct {
name string
patterns []string
requestOrigin string
expected []string
}{
{
name: "wildcard match includes request origin",
patterns: []string{"*.junv.cc", "https://static.com"},
requestOrigin: "https://api.junv.cc",
expected: []string{"https://api.junv.cc", "https://static.com"},
},
{
name: "no wildcard match returns static origins",
patterns: []string{"*.junv.cc", "https://static.com"},
requestOrigin: "https://other.com",
expected: []string{"*.junv.cc", "https://static.com"},
},
{
name: "multiple wildcards, one matches",
patterns: []string{"*.junv.cc", "*.example.com"},
requestOrigin: "https://api.junv.cc",
expected: []string{"https://api.junv.cc"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
matcher := NewWildcardMatcher(tt.patterns)
result := matcher.GetAllowedOrigins(tt.requestOrigin, tt.patterns)
if len(result) != len(tt.expected) {
t.Errorf("GetAllowedOrigins() returned %d origins, expected %d", len(result), len(tt.expected))
return
}
for i, expected := range tt.expected {
if result[i] != expected {
t.Errorf("GetAllowedOrigins()[%d] = %v, expected %v", i, result[i], expected)
}
}
})
}
}
-118
View File
@@ -1,118 +0,0 @@
apiVersion: v1
kind: ConfigMap
metadata:
name: passkey-auth-config
namespace: passkey-auth
data:
config.yaml: |
server:
port: "8080"
host: "0.0.0.0"
webauthn:
rp_display_name: "Passkey Auth"
rp_id: "your-domain.com"
rp_origins:
- "https://your-domain.com"
database:
path: "/data/passkey-auth.db"
cors:
allowed_origins:
- "https://your-domain.com"
auth:
session_secret: "your-session-secret-change-me"
require_approval: true
---
apiVersion: v1
kind: PersistentVolumeClaim
metadata:
name: passkey-auth-storage
namespace: passkey-auth
spec:
accessModes:
- ReadWriteOnce
resources:
requests:
storage: 1Gi
---
apiVersion: apps/v1
kind: Deployment
metadata:
name: passkey-auth
namespace: passkey-auth
labels:
app: passkey-auth
spec:
replicas: 1
selector:
matchLabels:
app: passkey-auth
template:
metadata:
labels:
app: passkey-auth
spec:
containers:
- name: passkey-auth
image: passkey-auth:latest
ports:
- containerPort: 8080
env:
- name: CONFIG_PATH
value: "/config/config.yaml"
- name: DATABASE_PATH
value: "/data/passkey-auth.db"
volumeMounts:
- name: config
mountPath: /config
readOnly: true
- name: data
mountPath: /data
livenessProbe:
httpGet:
path: /health
port: 8080
initialDelaySeconds: 30
periodSeconds: 30
readinessProbe:
httpGet:
path: /health
port: 8080
initialDelaySeconds: 5
periodSeconds: 5
resources:
requests:
memory: "64Mi"
cpu: "50m"
limits:
memory: "128Mi"
cpu: "100m"
volumes:
- name: config
configMap:
name: passkey-auth-config
- name: data
persistentVolumeClaim:
claimName: passkey-auth-storage
---
apiVersion: v1
kind: Service
metadata:
name: passkey-auth-service
namespace: passkey-auth
labels:
app: passkey-auth
spec:
selector:
app: passkey-auth
ports:
- port: 80
targetPort: 8080
protocol: TCP
type: ClusterIP
-44
View File
@@ -1,44 +0,0 @@
apiVersion: networking.k8s.io/v1
kind: Ingress
metadata:
name: example-app-ingress
namespace: default
annotations:
nginx.ingress.kubernetes.io/auth-url: "http://passkey-auth-service.passkey-auth.svc.cluster.local/auth"
nginx.ingress.kubernetes.io/auth-signin: "https://your-domain.com/auth"
nginx.ingress.kubernetes.io/auth-response-headers: "X-Auth-User,X-Auth-User-ID"
spec:
ingressClassName: nginx
rules:
- host: your-app.com
http:
paths:
- path: /
pathType: Prefix
backend:
service:
name: your-app-service
port:
number: 80
---
apiVersion: networking.k8s.io/v1
kind: Ingress
metadata:
name: passkey-auth-ingress
namespace: passkey-auth
annotations:
nginx.ingress.kubernetes.io/rewrite-target: /
spec:
ingressClassName: nginx
rules:
- host: your-domain.com
http:
paths:
- path: /auth
pathType: Prefix
backend:
service:
name: passkey-auth-service
port:
number: 80
-4
View File
@@ -1,4 +0,0 @@
apiVersion: v1
kind: Namespace
metadata:
name: passkey-auth
+4 -4
View File
@@ -8,11 +8,11 @@ import (
"time"
"github.com/gorilla/mux"
"github.com/rs/cors"
"github.com/sirupsen/logrus"
"passkey-auth/internal/auth"
"passkey-auth/internal/config"
"passkey-auth/internal/cors"
"passkey-auth/internal/database"
"passkey-auth/internal/handlers"
)
@@ -74,15 +74,15 @@ func main() {
// Static files for admin UI
router.PathPrefix("/").Handler(http.FileServer(http.Dir("./web/"))).Methods("GET")
// Setup CORS
c := cors.New(cors.Options{
// Setup CORS with wildcard support
corsHandler := cors.WildcardCORS(cors.Config{
AllowedOrigins: cfg.CORS.AllowedOrigins,
AllowedMethods: []string{"GET", "POST", "PUT", "DELETE", "OPTIONS"},
AllowedHeaders: []string{"*"},
AllowCredentials: true,
})
handler := c.Handler(router)
handler := corsHandler(router)
// Start server
port := os.Getenv("PORT")