331 lines
9.7 KiB
Go
331 lines
9.7 KiB
Go
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
|
|
}
|