package main import ( "encoding/json" "fmt" "log" "os" "strconv" "strings" "time" "unicode" "github.com/nixc/reverse-ssh-traefik/internal/buildinfo" "github.com/nixc/reverse-ssh-traefik/internal/client" "github.com/nixc/reverse-ssh-traefik/internal/sshutil" ) func envRequired(key string) string { v := os.Getenv(key) if v == "" { log.Fatalf("Required environment variable %s is not set", key) } return v } func envOr(key, fallback string) string { if v := os.Getenv(key); v != "" { return v } return fallback } func envBool(key string, fallback bool) bool { v := os.Getenv(key) if v == "" { return fallback } b, err := strconv.ParseBool(v) if err != nil { log.Printf("WARN: invalid %s=%q, using default %t", key, v, fallback) return fallback } return b } func parseLabels(raw string) (map[string]string, error) { raw = strings.TrimSpace(raw) if raw == "" { return nil, nil } if strings.HasPrefix(raw, "{") { labels := make(map[string]string) if err := json.Unmarshal([]byte(raw), &labels); err != nil { return nil, fmt.Errorf("parse JSON labels: %w", err) } return cleanLabels(labels), nil } labels := make(map[string]string) for _, line := range strings.FieldsFunc(raw, func(r rune) bool { return r == '\n' || r == ';' }) { line = strings.TrimSpace(line) if line == "" || strings.HasPrefix(line, "#") { continue } key, value, ok := strings.Cut(line, "=") if !ok { return nil, fmt.Errorf("label %q must be key=value", line) } key = strings.TrimSpace(key) value = strings.TrimSpace(value) if key == "" { return nil, fmt.Errorf("label %q has an empty key", line) } labels[key] = value } return labels, nil } func cleanLabels(labels map[string]string) map[string]string { cleaned := make(map[string]string, len(labels)) for key, value := range labels { cleanKey := strings.TrimSpace(key) cleanValue := strings.TrimSpace(value) if cleanKey != "" { cleaned[cleanKey] = cleanValue } } if len(cleaned) == 0 { return nil } return cleaned } func loadCustomLabels() (map[string]string, error) { labels := make(map[string]string) if path := cleanHost(os.Getenv("TUNNEL_LABELS_FILE")); path != "" { // #nosec G304 G703 -- operator-supplied local config file path. data, err := os.ReadFile(path) if err != nil { return nil, fmt.Errorf("read TUNNEL_LABELS_FILE %q: %w", path, err) } fileLabels, err := parseLabels(string(data)) if err != nil { return nil, fmt.Errorf("parse TUNNEL_LABELS_FILE %q: %w", path, err) } for key, value := range fileLabels { labels[key] = value } } envLabels, err := parseLabels(os.Getenv("TUNNEL_LABELS")) if err != nil { return nil, fmt.Errorf("parse TUNNEL_LABELS: %w", err) } for key, value := range envLabels { labels[key] = value } if len(labels) == 0 { return nil, nil } return labels, nil } // cleanHost strips BOM, NULs, and all Unicode whitespace (including NBSP) from ends. func cleanHost(s string) string { s = strings.ReplaceAll(s, "\x00", "") s = strings.TrimPrefix(s, "\ufeff") s = strings.TrimFunc(s, unicode.IsSpace) return s } // resolveBackendHost picks backend hostname. Priority: // 1. TUNNEL_HOST_FILE (path to a file containing hostname, e.g. Docker secret mount) // 2. TUNNEL_HOST // 3. TUNNEL_BACKEND_HOST (alias for stacks that reserve TUNNEL_HOST) // 4. 127.0.0.1 func resolveBackendHost() (host string, source string) { path := cleanHost(os.Getenv("TUNNEL_HOST_FILE")) if path != "" { // #nosec G304 G703 -- operator-supplied local backend host file path. b, err := os.ReadFile(path) if err != nil { log.Fatalf("TUNNEL_HOST_FILE %q: %v", path, err) } h := cleanHost(string(b)) if h != "" { return h, fmt.Sprintf("TUNNEL_HOST_FILE(%s)", path) } log.Fatalf("TUNNEL_HOST_FILE %q is empty after trim", path) } for _, key := range []string{"TUNNEL_HOST", "TUNNEL_BACKEND_HOST"} { raw := os.Getenv(key) h := cleanHost(raw) if h != "" { return h, key } if raw != "" && h == "" { log.Printf("tunnel-client: %s is set but only whitespace/control after clean (len=%d); trying next key / default", key, len(raw)) } } return "127.0.0.1", "default(127.0.0.1)" } func main() { log.SetFlags(log.LstdFlags | log.Lshortfile) log.Printf("tunnel-client starting commit=%s", buildinfo.Commit) serverAddr := envRequired("TUNNEL_SERVER") domain := envRequired("TUNNEL_DOMAIN") keyPath := envOr("TUNNEL_KEY", "/keys/id_ed25519") knownHostsPath := envOr("TUNNEL_KNOWN_HOSTS", "") if knownHostsPath == "" { knownHostsPath = sshutil.ExistingPath("/keys/known_hosts") } strictHostKey := envBool("TUNNEL_STRICT_HOST_KEY", false) hostKeyCallback, hostKeyMode, err := sshutil.HostKeyCallback(knownHostsPath, strictHostKey, "tunnel-server") if err != nil { log.Fatalf("Invalid tunnel server host key config: %v", err) } log.Printf("tunnel-server host key mode: %s", hostKeyMode) // Optional HTTP Basic Auth credentials for Traefik middleware. authUser := envOr("TUNNEL_AUTH_USER", "") authPass := envOr("TUNNEL_AUTH_PASS", "") customLabels, err := loadCustomLabels() if err != nil { log.Fatalf("Invalid custom labels: %v", err) } if len(customLabels) > 0 { log.Printf("Loaded %d custom tunnel label(s)", len(customLabels)) } localHost, hostSrc := resolveBackendHost() localPortStr := cleanHost(os.Getenv("TUNNEL_PORT")) if localPortStr == "" { localPortStr = "8080" } localPort, err := strconv.Atoi(localPortStr) if err != nil { log.Fatalf("Invalid TUNNEL_PORT=%q: %v", localPortStr, err) } log.Printf("tunnel-client config: backend host %q (source=%s); port=%d (TUNNEL_PORT raw=%q)", localHost, hostSrc, localPort, os.Getenv("TUNNEL_PORT")) log.Printf("tunnel-client env (for debugging): TUNNEL_HOST=%q TUNNEL_BACKEND_HOST=%q TUNNEL_HOST_FILE=%q", os.Getenv("TUNNEL_HOST"), os.Getenv("TUNNEL_BACKEND_HOST"), os.Getenv("TUNNEL_HOST_FILE")) // Load the private key. signer, err := client.LoadPrivateKey(keyPath) if err != nil { log.Fatalf("Failed to load private key: %v", err) } log.Printf("Loaded key from %s", keyPath) // Reconnect loop. backoff := time.Second maxBackoff := 30 * time.Second for { if authUser != "" { log.Printf("Connecting to %s (domain=%s, local=%s:%d, basicauth=enabled)", serverAddr, domain, localHost, localPort) } else { log.Printf("Connecting to %s (domain=%s, local=%s:%d)", serverAddr, domain, localHost, localPort) } sshClient, err := client.Connect(serverAddr, signer, hostKeyCallback) if err != nil { log.Printf("Connection failed: %v (retry in %s)", err, backoff) time.Sleep(backoff) backoff = min(backoff*2, maxBackoff) continue } // Reset backoff on successful connection. backoff = time.Second log.Printf("Connected to %s", serverAddr) // Set up the reverse tunnel (blocks until disconnected). if err := client.SetupTunnel(sshClient, domain, localHost, localPort, authUser, authPass, customLabels); err != nil { log.Printf("Tunnel error: %v (reconnecting in %s)", err, backoff) } sshClient.Close() time.Sleep(backoff) backoff = min(backoff*2, maxBackoff) } }