196 lines
5.1 KiB
Go
196 lines
5.1 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"crypto/hmac"
|
|
"crypto/sha256"
|
|
"crypto/tls"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/url"
|
|
"os"
|
|
"os/signal"
|
|
"strings"
|
|
"syscall"
|
|
"time"
|
|
|
|
"forgejo.digital-droplets.de/philschlo/proxui/platform/config"
|
|
|
|
"github.com/coder/websocket"
|
|
)
|
|
|
|
func main() {
|
|
cfg, err := config.Load()
|
|
if err != nil {
|
|
slog.Error("failed to load config", "error", err)
|
|
os.Exit(1)
|
|
}
|
|
|
|
logger := slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: slog.LevelInfo}))
|
|
logger = logger.With("service", "console-proxy", "env", cfg.AppEnv)
|
|
|
|
signingKey := []byte(cfg.OperatorToken)
|
|
if len(signingKey) == 0 {
|
|
logger.Error("OPERATOR_TOKEN is required")
|
|
os.Exit(1)
|
|
}
|
|
|
|
mux := http.NewServeMux()
|
|
mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"status":"ok","service":"console-proxy"}`))
|
|
})
|
|
mux.HandleFunc("GET /ws", func(w http.ResponseWriter, r *http.Request) {
|
|
handleWebSocket(w, r, signingKey, logger)
|
|
})
|
|
|
|
server := &http.Server{
|
|
Addr: cfg.ConsoleProxyAddr,
|
|
Handler: mux,
|
|
ReadHeaderTimeout: 5 * time.Second,
|
|
}
|
|
|
|
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
|
defer stop()
|
|
|
|
go func() {
|
|
logger.Info("console proxy listening", "addr", cfg.ConsoleProxyAddr)
|
|
if err := server.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
|
logger.Error("console proxy failed", "error", err)
|
|
os.Exit(1)
|
|
}
|
|
}()
|
|
|
|
<-ctx.Done()
|
|
|
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
if err := server.Shutdown(shutdownCtx); err != nil {
|
|
logger.Error("console proxy shutdown failed", "error", err)
|
|
os.Exit(1)
|
|
}
|
|
}
|
|
|
|
func handleWebSocket(w http.ResponseWriter, r *http.Request, signingKey []byte, logger *slog.Logger) {
|
|
ticket := r.URL.Query().Get("ticket")
|
|
if ticket == "" {
|
|
http.Error(w, "missing ticket", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
payload, err := verifyProxyTicket(ticket, signingKey)
|
|
if err != nil {
|
|
logger.Error("invalid ticket", "error", err)
|
|
http.Error(w, "invalid ticket", http.StatusForbidden)
|
|
return
|
|
}
|
|
|
|
endpoint := strings.TrimRight(payload.Endpoint, "/")
|
|
proxmoxWSURL := strings.Replace(endpoint, "https://", "wss://", 1)
|
|
vncURL := fmt.Sprintf(
|
|
"%s/nodes/%s/qemu/%d/vncwebsocket?port=%d&vncticket=%s",
|
|
proxmoxWSURL,
|
|
url.PathEscape(payload.Node),
|
|
payload.VMID,
|
|
payload.VNC.Port,
|
|
url.QueryEscape(payload.VNC.Ticket),
|
|
)
|
|
|
|
proxmoxConn, _, err := websocket.Dial(r.Context(), vncURL, &websocket.DialOptions{
|
|
HTTPClient: &http.Client{
|
|
Transport: &http.Transport{
|
|
TLSClientConfig: &tls.Config{
|
|
InsecureSkipVerify: true,
|
|
},
|
|
},
|
|
},
|
|
})
|
|
if err != nil {
|
|
logger.Error("proxmox websocket dial failed", "error", err)
|
|
http.Error(w, "proxy connection failed", http.StatusBadGateway)
|
|
return
|
|
}
|
|
defer proxmoxConn.Close(websocket.StatusNormalClosure, "")
|
|
|
|
clientConn, err := websocket.Accept(w, r, &websocket.AcceptOptions{})
|
|
if err != nil {
|
|
logger.Error("client websocket accept failed", "error", err)
|
|
return
|
|
}
|
|
defer clientConn.Close(websocket.StatusNormalClosure, "")
|
|
|
|
ctx, cancel := context.WithCancel(r.Context())
|
|
defer cancel()
|
|
|
|
go proxyWebSocket(ctx, clientConn, proxmoxConn, logger)
|
|
proxyWebSocket(ctx, proxmoxConn, clientConn, logger)
|
|
}
|
|
|
|
func proxyWebSocket(ctx context.Context, dst *websocket.Conn, src *websocket.Conn, logger *slog.Logger) {
|
|
for {
|
|
_, msg, err := src.Read(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
if err := dst.Write(ctx, websocket.MessageBinary, msg); err != nil {
|
|
logger.Error("websocket write failed", "error", err)
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
type proxyTicketPayload struct {
|
|
ClusterID string `json:"cluster_id"`
|
|
Node string `json:"node"`
|
|
VMID int `json:"vmid"`
|
|
TenantID string `json:"tenant_id"`
|
|
Endpoint string `json:"endpoint"`
|
|
VNC vncTicket `json:"vnc"`
|
|
ExpiresAt time.Time `json:"expires_at"`
|
|
Signature string `json:"signature"`
|
|
}
|
|
|
|
type vncTicket struct {
|
|
Port int `json:"port"`
|
|
Ticket string `json:"ticket"`
|
|
User string `json:"user"`
|
|
Cert string `json:"cert"`
|
|
UPID string `json:"upid"`
|
|
}
|
|
|
|
func verifyProxyTicket(ticket string, signingKey []byte) (proxyTicketPayload, error) {
|
|
data, err := base64.RawURLEncoding.DecodeString(ticket)
|
|
if err != nil {
|
|
return proxyTicketPayload{}, fmt.Errorf("invalid ticket encoding")
|
|
}
|
|
|
|
sigData := make([]byte, len(data))
|
|
copy(sigData, data)
|
|
|
|
var payload proxyTicketPayload
|
|
if err := json.Unmarshal(data, &payload); err != nil {
|
|
return proxyTicketPayload{}, fmt.Errorf("invalid ticket payload")
|
|
}
|
|
|
|
if time.Now().After(payload.ExpiresAt) {
|
|
return proxyTicketPayload{}, fmt.Errorf("ticket expired")
|
|
}
|
|
|
|
receivedSig := payload.Signature
|
|
payload.Signature = ""
|
|
dataToVerify, _ := json.Marshal(payload)
|
|
|
|
mac := hmac.New(sha256.New, signingKey)
|
|
mac.Write(dataToVerify)
|
|
expectedSig := base64.RawURLEncoding.EncodeToString(mac.Sum(nil))
|
|
|
|
if !hmac.Equal([]byte(receivedSig), []byte(expectedSig)) {
|
|
return proxyTicketPayload{}, fmt.Errorf("invalid ticket signature")
|
|
}
|
|
|
|
return payload, nil
|
|
} |