548 lines
16 KiB
Go
548 lines
16 KiB
Go
package server
|
|
|
|
import (
|
|
"encoding/json"
|
|
"fmt"
|
|
"log"
|
|
"sort"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/nixc/reverse-ssh-traefik/internal/sshutil"
|
|
"golang.org/x/crypto/bcrypt"
|
|
"golang.org/x/crypto/ssh"
|
|
)
|
|
|
|
// LabelManager manages Traefik routing labels on its own Swarm service
|
|
// by SSHing into the Swarm manager and running docker service update.
|
|
type LabelManager struct {
|
|
mu sync.Mutex
|
|
remoteHost string // Swarm manager, e.g. "ingress.nixc.us"
|
|
remoteUser string // SSH user
|
|
signer ssh.Signer
|
|
hostKeyCB ssh.HostKeyCallback
|
|
serviceName string // Swarm service name, e.g. "better-argo-tunnels_tunnel-server"
|
|
entrypoint string // e.g. "websecure"
|
|
certResolver string // e.g. "letsencryptresolver"
|
|
labels map[string]bool // track which tunnel keys we've added
|
|
authLabels map[string]bool // track which tunnel keys have auth middleware
|
|
customLabels map[string][]string // track extra label keys added per tunnel
|
|
}
|
|
|
|
// NewLabelManager creates a label manager that updates Swarm service labels via SSH.
|
|
func NewLabelManager(
|
|
remoteHost, remoteUser string,
|
|
signer ssh.Signer,
|
|
hostKeyCB ssh.HostKeyCallback,
|
|
serviceName, entrypoint, certResolver string,
|
|
) (*LabelManager, error) {
|
|
|
|
lm := &LabelManager{
|
|
remoteHost: remoteHost,
|
|
remoteUser: remoteUser,
|
|
signer: signer,
|
|
hostKeyCB: hostKeyCB,
|
|
serviceName: serviceName,
|
|
entrypoint: entrypoint,
|
|
certResolver: certResolver,
|
|
labels: make(map[string]bool),
|
|
authLabels: make(map[string]bool),
|
|
customLabels: make(map[string][]string),
|
|
}
|
|
|
|
// Verify we can reach the Swarm manager and the service exists.
|
|
cmd := fmt.Sprintf("docker service inspect --format '{{.Spec.Name}}' %s", serviceName)
|
|
if err := lm.runRemote(cmd); err != nil {
|
|
log.Printf("WARN: could not verify service %s (may not exist yet): %v", serviceName, err)
|
|
} else {
|
|
log.Printf("Verified Swarm service: %s", serviceName)
|
|
}
|
|
|
|
log.Printf("Label manager ready (host=%s, service=%s, ep=%s, resolver=%s)",
|
|
remoteHost, serviceName, entrypoint, certResolver)
|
|
|
|
return lm, nil
|
|
}
|
|
|
|
// ReconcileExistingTunnelLabels adopts existing tunnel-* Traefik labels from
|
|
// the Swarm service so a restarted server preserves routes while clients
|
|
// reconnect. It does not remove labels; explicit purge remains opt-in.
|
|
func (lm *LabelManager) ReconcileExistingTunnelLabels() error {
|
|
labels, err := lm.inspectLabels()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
tunnelKeys, authKeys := collectExistingTunnelKeys(labels)
|
|
|
|
lm.mu.Lock()
|
|
for key := range tunnelKeys {
|
|
lm.labels[key] = true
|
|
}
|
|
for key := range authKeys {
|
|
lm.authLabels[key] = true
|
|
}
|
|
lm.mu.Unlock()
|
|
|
|
log.Printf("Reconciled %d existing tunnel route(s), %d with auth", len(tunnelKeys), len(authKeys))
|
|
return nil
|
|
}
|
|
|
|
// PurgeAllTunnelLabels removes every tunnel-* Traefik label from the Swarm
|
|
// service. This is intentionally opt-in because purging on every restart
|
|
// creates a 404 window until every client reconnects.
|
|
func (lm *LabelManager) PurgeAllTunnelLabels() error {
|
|
labels, err := lm.inspectLabels()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
var rmLabels []string
|
|
for key := range labels {
|
|
if isTunnelLabel(key) {
|
|
rmLabels = append(rmLabels, key)
|
|
}
|
|
}
|
|
|
|
if len(rmLabels) == 0 {
|
|
log.Println("No stale tunnel labels to purge")
|
|
return nil
|
|
}
|
|
|
|
rmCmd := fmt.Sprintf("docker service update --label-rm %s %s",
|
|
labelRmArgs(rmLabels), shellQuote(lm.serviceName))
|
|
if err := lm.runRemote(rmCmd); err != nil {
|
|
return fmt.Errorf("purge tunnel labels: %w", err)
|
|
}
|
|
|
|
log.Printf("Purged %d tunnel label(s)", len(rmLabels))
|
|
return nil
|
|
}
|
|
|
|
func (lm *LabelManager) inspectLabels() (map[string]string, error) {
|
|
addr := lm.remoteHost
|
|
if !strings.Contains(addr, ":") {
|
|
addr = addr + ":22"
|
|
}
|
|
|
|
config := &ssh.ClientConfig{
|
|
User: lm.remoteUser,
|
|
Auth: []ssh.AuthMethod{
|
|
ssh.PublicKeys(lm.signer),
|
|
},
|
|
Config: sshutil.SecureConfig(),
|
|
HostKeyAlgorithms: sshutil.HostKeyAlgorithms(),
|
|
HostKeyCallback: lm.hostKeyCB,
|
|
Timeout: 15 * time.Second,
|
|
}
|
|
|
|
client, err := ssh.Dial("tcp", addr, config)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("SSH dial %s: %w", addr, err)
|
|
}
|
|
defer client.Close()
|
|
|
|
session, err := client.NewSession()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("SSH session: %w", err)
|
|
}
|
|
defer session.Close()
|
|
|
|
inspectCmd := fmt.Sprintf(
|
|
"docker service inspect --format '{{json .Spec.Labels}}' %s", lm.serviceName)
|
|
output, err := session.CombinedOutput(inspectCmd)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("inspect labels: %w (%s)", err, string(output))
|
|
}
|
|
|
|
var labels map[string]string
|
|
if err := json.Unmarshal(output, &labels); err != nil {
|
|
return nil, fmt.Errorf("parse labels JSON: %w", err)
|
|
}
|
|
|
|
return labels, nil
|
|
}
|
|
|
|
func isTunnelLabel(key string) bool {
|
|
return strings.HasPrefix(key, "traefik.http.routers.tunnel-") ||
|
|
strings.HasPrefix(key, "traefik.http.services.tunnel-") ||
|
|
strings.HasPrefix(key, "traefik.http.middlewares.tunnel-")
|
|
}
|
|
|
|
func collectExistingTunnelKeys(labels map[string]string) (map[string]bool, map[string]bool) {
|
|
tunnelKeys := make(map[string]bool)
|
|
authKeys := make(map[string]bool)
|
|
|
|
for label := range labels {
|
|
if key, ok := tunnelKeyFromLabel(label, "traefik.http.routers.tunnel-", "-router."); ok {
|
|
tunnelKeys[key] = true
|
|
continue
|
|
}
|
|
if key, ok := tunnelKeyFromLabel(label, "traefik.http.services.tunnel-", "-service."); ok {
|
|
tunnelKeys[key] = true
|
|
continue
|
|
}
|
|
if key, ok := tunnelKeyFromLabel(label, "traefik.http.middlewares.tunnel-", "-auth."); ok {
|
|
tunnelKeys[key] = true
|
|
authKeys[key] = true
|
|
}
|
|
}
|
|
|
|
return tunnelKeys, authKeys
|
|
}
|
|
|
|
func tunnelKeyFromLabel(label, prefix, suffix string) (string, bool) {
|
|
if !strings.HasPrefix(label, prefix) {
|
|
return "", false
|
|
}
|
|
rest := strings.TrimPrefix(label, prefix)
|
|
idx := strings.Index(rest, suffix)
|
|
if idx <= 0 {
|
|
return "", false
|
|
}
|
|
return rest[:idx], true
|
|
}
|
|
|
|
// Add adds Traefik routing labels to the Swarm service for a tunnel.
|
|
// If authUser and authPass are non-empty, a basicauth middleware is also added.
|
|
func (lm *LabelManager) Add(
|
|
tunKey, domain string,
|
|
port int,
|
|
authUser, authPass string,
|
|
custom map[string]string,
|
|
) error {
|
|
lm.mu.Lock()
|
|
defer lm.mu.Unlock()
|
|
|
|
routerName := fmt.Sprintf("tunnel-%s-router", tunKey)
|
|
serviceName := fmt.Sprintf("tunnel-%s-service", tunKey)
|
|
middlewareName := fmt.Sprintf("tunnel-%s-auth", tunKey)
|
|
|
|
routerPrefix := fmt.Sprintf("traefik.http.routers.%s.", routerName)
|
|
servicePrefix := fmt.Sprintf("traefik.http.services.%s.", serviceName)
|
|
labels := map[string]string{
|
|
routerPrefix + "rule": fmt.Sprintf("Host(`%s`)", domain),
|
|
routerPrefix + "entrypoints": lm.entrypoint,
|
|
routerPrefix + "tls": "true",
|
|
routerPrefix + "tls.certresolver": lm.certResolver,
|
|
routerPrefix + "service": serviceName,
|
|
servicePrefix + "loadbalancer.server.port": fmt.Sprintf("%d", port),
|
|
}
|
|
customKeys := make([]string, 0, len(custom))
|
|
|
|
// If auth credentials are provided, add basicauth middleware labels.
|
|
authEnabled := authUser != "" && authPass != ""
|
|
if authEnabled {
|
|
htpasswd, err := generateHTPasswd(authUser, authPass)
|
|
if err != nil {
|
|
return fmt.Errorf("generate htpasswd for %s: %w", domain, err)
|
|
}
|
|
labels[fmt.Sprintf("traefik.http.middlewares.%s.basicauth.users", middlewareName)] = htpasswd
|
|
labels[routerPrefix+"middlewares"] = middlewareName
|
|
log.Printf("BasicAuth middleware %s added for %s", middlewareName, domain)
|
|
}
|
|
|
|
renderedCustom, err := renderCustomLabels(tunKey, domain, port, routerName, serviceName, middlewareName, custom)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
for key, value := range renderedCustom {
|
|
if !isSafeLabelKey(key) {
|
|
return fmt.Errorf("custom label %q contains unsupported characters", key)
|
|
}
|
|
if !isAllowedCustomLabel(key, tunKey, routerName, serviceName) {
|
|
return fmt.Errorf("custom label %q is outside tunnel %s", key, tunKey)
|
|
}
|
|
if isProtectedCustomLabel(key, routerName, serviceName) {
|
|
return fmt.Errorf("custom label %q would replace a managed tunnel label", key)
|
|
}
|
|
if key == routerPrefix+"middlewares" && labels[key] != "" {
|
|
value = mergeMiddlewares(labels[key], value)
|
|
}
|
|
labels[key] = value
|
|
customKeys = append(customKeys, key)
|
|
}
|
|
|
|
staleLabels := staleLabelKeys(lm.customLabels[tunKey], labels)
|
|
if lm.authLabels[tunKey] && !authEnabled {
|
|
staleLabels = append(staleLabels, fmt.Sprintf("traefik.http.middlewares.%s.basicauth.users", middlewareName))
|
|
if _, hasMiddlewares := labels[routerPrefix+"middlewares"]; !hasMiddlewares {
|
|
staleLabels = append(staleLabels, routerPrefix+"middlewares")
|
|
}
|
|
}
|
|
staleLabels = uniqueStrings(staleLabels)
|
|
sort.Strings(staleLabels)
|
|
|
|
labelArgs := make([]string, 0, len(labels))
|
|
for _, key := range sortedKeys(labels) {
|
|
labelArgs = append(labelArgs, labelFlag(key, labels[key]))
|
|
}
|
|
|
|
cmdParts := []string{"docker service update"}
|
|
for _, key := range staleLabels {
|
|
cmdParts = append(cmdParts, "--label-rm "+shellQuote(key))
|
|
}
|
|
for _, arg := range labelArgs {
|
|
cmdParts = append(cmdParts, "--label-add "+arg)
|
|
}
|
|
cmdParts = append(cmdParts, shellQuote(lm.serviceName))
|
|
cmd := strings.Join(cmdParts, " ")
|
|
|
|
if err := lm.runRemote(cmd); err != nil {
|
|
return fmt.Errorf("add labels for %s: %w", domain, err)
|
|
}
|
|
|
|
lm.labels[tunKey] = true
|
|
if authEnabled {
|
|
lm.authLabels[tunKey] = true
|
|
} else {
|
|
delete(lm.authLabels, tunKey)
|
|
}
|
|
if len(customKeys) > 0 {
|
|
sort.Strings(customKeys)
|
|
lm.customLabels[tunKey] = customKeys
|
|
log.Printf("Added %d custom label(s) for %s", len(customKeys), domain)
|
|
} else {
|
|
delete(lm.customLabels, tunKey)
|
|
}
|
|
log.Printf("Added Swarm labels: %s -> %s:%d", domain, lm.serviceName, port)
|
|
return nil
|
|
}
|
|
|
|
// Remove removes Traefik routing labels from the Swarm service for a tunnel.
|
|
func (lm *LabelManager) Remove(tunKey string) error {
|
|
lm.mu.Lock()
|
|
defer lm.mu.Unlock()
|
|
|
|
if !lm.labels[tunKey] {
|
|
return nil // nothing to remove
|
|
}
|
|
|
|
routerName := fmt.Sprintf("tunnel-%s-router", tunKey)
|
|
serviceName := fmt.Sprintf("tunnel-%s-service", tunKey)
|
|
|
|
middlewareName := fmt.Sprintf("tunnel-%s-auth", tunKey)
|
|
|
|
// Build the label-rm flags.
|
|
rmLabels := []string{
|
|
fmt.Sprintf("traefik.http.routers.%s.rule", routerName),
|
|
fmt.Sprintf("traefik.http.routers.%s.entrypoints", routerName),
|
|
fmt.Sprintf("traefik.http.routers.%s.tls", routerName),
|
|
fmt.Sprintf("traefik.http.routers.%s.tls.certresolver", routerName),
|
|
fmt.Sprintf("traefik.http.routers.%s.service", routerName),
|
|
fmt.Sprintf("traefik.http.services.%s.loadbalancer.server.port", serviceName),
|
|
}
|
|
|
|
// Remove auth middleware labels if they were added.
|
|
if lm.authLabels[tunKey] {
|
|
rmLabels = append(rmLabels,
|
|
fmt.Sprintf("traefik.http.middlewares.%s.basicauth.users", middlewareName),
|
|
fmt.Sprintf("traefik.http.routers.%s.middlewares", routerName),
|
|
)
|
|
delete(lm.authLabels, tunKey)
|
|
log.Printf("Removing BasicAuth middleware %s", middlewareName)
|
|
}
|
|
rmLabels = append(rmLabels, lm.customLabels[tunKey]...)
|
|
delete(lm.customLabels, tunKey)
|
|
rmLabels = uniqueStrings(rmLabels)
|
|
|
|
cmd := fmt.Sprintf("docker service update --label-rm %s %s",
|
|
labelRmArgs(rmLabels), shellQuote(lm.serviceName))
|
|
|
|
if err := lm.runRemote(cmd); err != nil {
|
|
return fmt.Errorf("remove labels for %s: %w", tunKey, err)
|
|
}
|
|
|
|
delete(lm.labels, tunKey)
|
|
log.Printf("Removed Swarm labels for tunnel: %s", tunKey)
|
|
return nil
|
|
}
|
|
|
|
// generateHTPasswd creates a bcrypt-hashed htpasswd entry for Traefik basicauth.
|
|
// The output format is user:$hash. Dollar signs are NOT doubled here because
|
|
// we pass labels via docker service update with single-quoted values, which
|
|
// preserves them literally. Doubling is only needed in compose files.
|
|
func generateHTPasswd(user, pass string) (string, error) {
|
|
hash, err := bcrypt.GenerateFromPassword([]byte(pass), bcrypt.DefaultCost)
|
|
if err != nil {
|
|
return "", fmt.Errorf("bcrypt hash: %w", err)
|
|
}
|
|
return fmt.Sprintf("%s:%s", user, string(hash)), nil
|
|
}
|
|
|
|
// labelFlag formats a --label-add value, quoting properly for shell.
|
|
func labelFlag(key, value string) string {
|
|
return shellQuote(fmt.Sprintf("%s=%s", key, value))
|
|
}
|
|
|
|
func shellQuote(value string) string {
|
|
return "'" + strings.ReplaceAll(value, "'", "'\\''") + "'"
|
|
}
|
|
|
|
func labelRmArgs(keys []string) string {
|
|
args := make([]string, 0, len(keys))
|
|
for _, key := range keys {
|
|
args = append(args, shellQuote(key))
|
|
}
|
|
return strings.Join(args, " --label-rm ")
|
|
}
|
|
|
|
func renderCustomLabels(
|
|
tunKey, domain string,
|
|
port int,
|
|
routerName, serviceName, middlewareName string,
|
|
custom map[string]string,
|
|
) (map[string]string, error) {
|
|
if len(custom) == 0 {
|
|
return nil, nil
|
|
}
|
|
|
|
replacer := strings.NewReplacer(
|
|
"{tunKey}", tunKey,
|
|
"{domain}", domain,
|
|
"{port}", fmt.Sprintf("%d", port),
|
|
"{router}", routerName,
|
|
"{service}", serviceName,
|
|
"{middleware}", middlewareName,
|
|
)
|
|
rendered := make(map[string]string, len(custom))
|
|
for rawKey, rawValue := range custom {
|
|
key := strings.TrimSpace(replacer.Replace(rawKey))
|
|
if key == "" {
|
|
return nil, fmt.Errorf("custom label has an empty key")
|
|
}
|
|
rendered[key] = strings.TrimSpace(replacer.Replace(rawValue))
|
|
}
|
|
return rendered, nil
|
|
}
|
|
|
|
func isAllowedCustomLabel(key, tunKey, routerName, serviceName string) bool {
|
|
return strings.HasPrefix(key, fmt.Sprintf("traefik.http.routers.%s.", routerName)) ||
|
|
strings.HasPrefix(key, fmt.Sprintf("traefik.http.services.%s.", serviceName)) ||
|
|
strings.HasPrefix(key, fmt.Sprintf("traefik.http.middlewares.tunnel-%s-", tunKey))
|
|
}
|
|
|
|
func isProtectedCustomLabel(key, routerName, serviceName string) bool {
|
|
protected := map[string]bool{
|
|
fmt.Sprintf("traefik.http.routers.%s.rule", routerName): true,
|
|
fmt.Sprintf("traefik.http.routers.%s.entrypoints", routerName): true,
|
|
fmt.Sprintf("traefik.http.routers.%s.tls", routerName): true,
|
|
fmt.Sprintf("traefik.http.routers.%s.tls.certresolver", routerName): true,
|
|
fmt.Sprintf("traefik.http.routers.%s.service", routerName): true,
|
|
fmt.Sprintf("traefik.http.services.%s.loadbalancer.server.port", serviceName): true,
|
|
}
|
|
return protected[key]
|
|
}
|
|
|
|
func isSafeLabelKey(key string) bool {
|
|
for _, r := range key {
|
|
if r >= 'a' && r <= 'z' {
|
|
continue
|
|
}
|
|
if r >= 'A' && r <= 'Z' {
|
|
continue
|
|
}
|
|
if r >= '0' && r <= '9' {
|
|
continue
|
|
}
|
|
if r == '.' || r == '-' || r == '_' {
|
|
continue
|
|
}
|
|
return false
|
|
}
|
|
return key != ""
|
|
}
|
|
|
|
func mergeMiddlewares(existing, extra string) string {
|
|
seen := make(map[string]bool)
|
|
var merged []string
|
|
for _, list := range []string{existing, extra} {
|
|
for _, item := range strings.Split(list, ",") {
|
|
item = strings.TrimSpace(item)
|
|
if item == "" || seen[item] {
|
|
continue
|
|
}
|
|
seen[item] = true
|
|
merged = append(merged, item)
|
|
}
|
|
}
|
|
return strings.Join(merged, ",")
|
|
}
|
|
|
|
func sortedKeys(values map[string]string) []string {
|
|
keys := make([]string, 0, len(values))
|
|
for key := range values {
|
|
keys = append(keys, key)
|
|
}
|
|
sort.Strings(keys)
|
|
return keys
|
|
}
|
|
|
|
func uniqueStrings(values []string) []string {
|
|
seen := make(map[string]bool, len(values))
|
|
unique := make([]string, 0, len(values))
|
|
for _, value := range values {
|
|
if seen[value] {
|
|
continue
|
|
}
|
|
seen[value] = true
|
|
unique = append(unique, value)
|
|
}
|
|
return unique
|
|
}
|
|
|
|
func staleLabelKeys(oldKeys []string, currentLabels map[string]string) []string {
|
|
var stale []string
|
|
for _, key := range oldKeys {
|
|
if _, stillPresent := currentLabels[key]; !stillPresent {
|
|
stale = append(stale, key)
|
|
}
|
|
}
|
|
sort.Strings(stale)
|
|
return stale
|
|
}
|
|
|
|
// runRemote executes a command on the Swarm manager via SSH.
|
|
func (lm *LabelManager) runRemote(cmd string) error {
|
|
addr := lm.remoteHost
|
|
if !strings.Contains(addr, ":") {
|
|
addr = addr + ":22"
|
|
}
|
|
|
|
config := &ssh.ClientConfig{
|
|
User: lm.remoteUser,
|
|
Auth: []ssh.AuthMethod{
|
|
ssh.PublicKeys(lm.signer),
|
|
},
|
|
Config: sshutil.SecureConfig(),
|
|
HostKeyAlgorithms: sshutil.HostKeyAlgorithms(),
|
|
HostKeyCallback: lm.hostKeyCB,
|
|
Timeout: 15 * time.Second,
|
|
}
|
|
|
|
client, err := ssh.Dial("tcp", addr, config)
|
|
if err != nil {
|
|
return fmt.Errorf("SSH dial %s: %w", addr, err)
|
|
}
|
|
defer client.Close()
|
|
|
|
session, err := client.NewSession()
|
|
if err != nil {
|
|
return fmt.Errorf("SSH session: %w", err)
|
|
}
|
|
defer session.Close()
|
|
|
|
output, err := session.CombinedOutput(cmd)
|
|
if err != nil {
|
|
return fmt.Errorf("remote cmd failed: %w (output: %s)", err, string(output))
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// Close is a no-op — SSH connections are opened/closed per operation.
|
|
func (lm *LabelManager) Close() error {
|
|
return nil
|
|
}
|