252 lines
7.0 KiB
Go
252 lines
7.0 KiB
Go
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)
|
|
}
|
|
}
|