package config import ( "crypto/ed25519" "encoding/base64" "encoding/json" "errors" "fmt" "net/netip" "net/url" "os" "path/filepath" "strconv" "strings" "time" "github.com/klarkxy/nekonest-cloud/relay/internal/authsnapshot" ) type Config struct { ListenAddress string DataRoot string BackupRoot string NodeID string ControlPlaneURL string ClientCertificate string ClientKey string ControlPlaneCA string InternalRelayCA string InternalEndpoints map[string]string AllowedPWAOrigins []string RouteSecret []byte SourceHashSecret []byte ForwardSecret []byte HandoffSecret []byte SnapshotKeys authsnapshot.Keyring TrustedProxyRanges []netip.Prefix MaxTenants int ShutdownTimeout time.Duration } func requiredEnv(name string) (string, error) { value := strings.TrimSpace(os.Getenv(name)) if value == "" { return "", fmt.Errorf("%s is required", name) } return value, nil } func decodeSecret(name string) ([]byte, error) { raw, err := requiredEnv(name) if err != nil { return nil, err } decoded, decodeErr := base64.RawURLEncoding.DecodeString(raw) if decodeErr != nil { decoded, decodeErr = base64.StdEncoding.DecodeString(raw) } if decodeErr != nil || len(decoded) < 32 { return nil, fmt.Errorf("%s must be at least 32 random bytes encoded as base64url", name) } return decoded, nil } func exactOrigin(raw string) (string, error) { parsed, err := url.Parse(strings.TrimSpace(raw)) if err != nil || parsed.Host == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { return "", errors.New("invalid origin") } loopback := parsed.Hostname() == "localhost" || parsed.Hostname() == "127.0.0.1" || parsed.Hostname() == "::1" if parsed.Scheme != "https" && !(parsed.Scheme == "http" && loopback) { return "", errors.New("origin must use HTTPS") } if parsed.Path != "" && parsed.Path != "/" { return "", errors.New("origin must not contain a path") } return parsed.Scheme + "://" + parsed.Host, nil } func parseOrigins(raw string) ([]string, error) { seen := make(map[string]struct{}) var origins []string for _, item := range strings.Split(raw, ",") { if strings.TrimSpace(item) == "" { continue } origin, err := exactOrigin(item) if err != nil { return nil, fmt.Errorf("invalid PWA origin %q: %w", item, err) } if _, exists := seen[origin]; exists { continue } seen[origin] = struct{}{} origins = append(origins, origin) } if len(origins) == 0 { return nil, errors.New("NEKONEST_RELAY_PWA_ORIGINS requires at least one exact origin") } return origins, nil } type snapshotKeyConfig struct { KID string `json:"kid"` PublicKeyJWK json.RawMessage `json:"public_key_jwk"` NotBefore string `json:"not_before"` RetainUntil string `json:"retain_until"` } func parseSnapshotKeys(raw string, now time.Time) (authsnapshot.Keyring, error) { var records []snapshotKeyConfig decoder := json.NewDecoder(strings.NewReader(raw)) decoder.DisallowUnknownFields() if err := decoder.Decode(&records); err != nil || len(records) == 0 { return nil, errors.New("NEKONEST_RELAY_SNAPSHOT_KEYS must be a non-empty JSON array") } keyring := make(authsnapshot.Keyring, len(records)) for _, record := range records { if record.KID == "" { return nil, errors.New("snapshot key kid is required") } key, jwkKID, err := authsnapshot.PublicKeyFromJWK(record.PublicKeyJWK) if err != nil || (jwkKID != "" && jwkKID != record.KID) { return nil, fmt.Errorf("invalid snapshot key %q", record.KID) } notBefore, err := time.Parse(time.RFC3339Nano, record.NotBefore) if err != nil { return nil, fmt.Errorf("invalid not_before for snapshot key %q", record.KID) } retainUntil, err := time.Parse(time.RFC3339Nano, record.RetainUntil) if err != nil || !retainUntil.After(notBefore) { return nil, fmt.Errorf("invalid retain_until for snapshot key %q", record.KID) } if !retainUntil.After(now) { continue } keyring[record.KID] = authsnapshot.Key{ PublicKey: ed25519.PublicKey(append([]byte(nil), key...)), NotBefore: notBefore, RetainUntil: retainUntil, } } if len(keyring) == 0 { return nil, errors.New("no unexpired snapshot verification key is configured") } return keyring, nil } func parseProxyRanges(raw string) ([]netip.Prefix, error) { var prefixes []netip.Prefix for _, item := range strings.Split(raw, ",") { item = strings.TrimSpace(item) if item == "" { continue } prefix, err := netip.ParsePrefix(item) if err != nil { return nil, fmt.Errorf("invalid trusted proxy range %q", item) } prefixes = append(prefixes, prefix.Masked()) } return prefixes, nil } func parseInternalEndpoints(raw string) (map[string]string, error) { if strings.TrimSpace(raw) == "" { return map[string]string{}, nil } var values map[string]string decoder := json.NewDecoder(strings.NewReader(raw)) decoder.DisallowUnknownFields() if err := decoder.Decode(&values); err != nil || values == nil { return nil, errors.New("NEKONEST_RELAY_INTERNAL_ENDPOINTS must be a JSON object") } result := make(map[string]string, len(values)) for reference, rawOrigin := range values { if !validEndpointReference(reference) { return nil, fmt.Errorf("invalid internal endpoint reference %q", reference) } origin, err := exactOrigin(rawOrigin) if err != nil { return nil, fmt.Errorf("invalid internal endpoint %q: %w", reference, err) } result[reference] = origin } return result, nil } func validEndpointReference(value string) bool { if value == "" || len(value) > 128 { return false } for index, char := range value { if (char >= 'A' && char <= 'Z') || (char >= 'a' && char <= 'z') || (char >= '0' && char <= '9') || (index > 0 && strings.ContainsRune("._:-", char)) { continue } return false } return true } func parsePositiveInt(name string, fallback int) (int, error) { raw := strings.TrimSpace(os.Getenv(name)) if raw == "" { return fallback, nil } value, err := strconv.Atoi(raw) if err != nil || value <= 0 { return 0, fmt.Errorf("%s must be a positive integer", name) } return value, nil } func Load() (Config, error) { dataRoot, err := requiredEnv("NEKONEST_RELAY_DATA_ROOT") if err != nil { return Config{}, err } dataRoot, err = filepath.Abs(dataRoot) if err != nil { return Config{}, fmt.Errorf("resolve relay data root: %w", err) } backupRoot, err := requiredEnv("NEKONEST_RELAY_BACKUP_ROOT") if err != nil { return Config{}, err } backupRoot, err = filepath.Abs(backupRoot) if err != nil { return Config{}, fmt.Errorf("resolve relay backup root: %w", err) } for _, pair := range [][2]string{{dataRoot, backupRoot}, {backupRoot, dataRoot}} { relative, relErr := filepath.Rel(pair[0], pair[1]) if relErr == nil && (relative == "." || (relative != ".." && !strings.HasPrefix(relative, ".."+string(os.PathSeparator)))) { return Config{}, errors.New("NEKONEST_RELAY_DATA_ROOT and NEKONEST_RELAY_BACKUP_ROOT must be separate directory trees") } } nodeID, err := requiredEnv("NEKONEST_RELAY_NODE_ID") if err != nil || !strings.HasPrefix(nodeID, "node_") { return Config{}, errors.New("NEKONEST_RELAY_NODE_ID must be a node_ identity") } controlPlaneURL, err := requiredEnv("NEKONEST_RELAY_CONTROL_PLANE_URL") if err != nil { return Config{}, err } if _, err := exactOrigin(controlPlaneURL); err != nil { return Config{}, fmt.Errorf("invalid control plane URL: %w", err) } certificate, err := requiredEnv("NEKONEST_RELAY_MTLS_CERT_FILE") if err != nil { return Config{}, err } key, err := requiredEnv("NEKONEST_RELAY_MTLS_KEY_FILE") if err != nil { return Config{}, err } ca, err := requiredEnv("NEKONEST_RELAY_CONTROL_PLANE_CA_FILE") if err != nil { return Config{}, err } internalCA, err := requiredEnv("NEKONEST_RELAY_INTERNAL_CA_FILE") if err != nil { return Config{}, err } internalEndpoints, err := parseInternalEndpoints(os.Getenv("NEKONEST_RELAY_INTERNAL_ENDPOINTS")) if err != nil { return Config{}, err } origins, err := parseOrigins(os.Getenv("NEKONEST_RELAY_PWA_ORIGINS")) if err != nil { return Config{}, err } routeSecret, err := decodeSecret("NEKONEST_RELAY_ROUTE_SECRET") if err != nil { return Config{}, err } sourceSecret, err := decodeSecret("NEKONEST_RELAY_SOURCE_HASH_SECRET") if err != nil { return Config{}, err } forwardSecret, err := decodeSecret("NEKONEST_RELAY_FORWARD_SECRET") if err != nil { return Config{}, err } handoffSecret, err := decodeSecret("NEKONEST_RELAY_HANDOFF_SECRET") if err != nil { return Config{}, err } keyJSON, err := requiredEnv("NEKONEST_RELAY_SNAPSHOT_KEYS") if err != nil { return Config{}, err } keyring, err := parseSnapshotKeys(keyJSON, time.Now()) if err != nil { return Config{}, err } proxyRanges, err := parseProxyRanges(os.Getenv("NEKONEST_RELAY_TRUSTED_PROXY_CIDRS")) if err != nil { return Config{}, err } maxTenants, err := parsePositiveInt("NEKONEST_RELAY_MAX_TENANTS", 256) if err != nil { return Config{}, err } listen := strings.TrimSpace(os.Getenv("NEKONEST_RELAY_LISTEN")) if listen == "" { listen = ":8080" } return Config{ ListenAddress: listen, DataRoot: dataRoot, BackupRoot: backupRoot, NodeID: nodeID, ControlPlaneURL: controlPlaneURL, ClientCertificate: certificate, ClientKey: key, ControlPlaneCA: ca, InternalRelayCA: internalCA, InternalEndpoints: internalEndpoints, AllowedPWAOrigins: origins, RouteSecret: routeSecret, SourceHashSecret: sourceSecret, ForwardSecret: forwardSecret, HandoffSecret: handoffSecret, SnapshotKeys: keyring, TrustedProxyRanges: proxyRanges, MaxTenants: maxTenants, ShutdownTimeout: 15 * time.Second, }, nil }