Files
nekonest-cloud/relay/internal/config/config.go
T

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
}