feat: share proxmox cluster client with worker

This commit is contained in:
Philipp
2026-06-11 10:24:36 +02:00
parent 7c818cf92e
commit bf31a8db37
16 changed files with 240 additions and 14 deletions
+242
View File
@@ -0,0 +1,242 @@
package proxmox
import (
"bytes"
"context"
"crypto/sha256"
"crypto/subtle"
"crypto/tls"
"crypto/x509"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
"forgejo.digital-droplets.de/philschlo/proxui/platform/cluster"
)
const defaultTimeout = 15 * time.Second
type Client struct {
baseURL *url.URL
tokenID string
tokenSecret string
httpClient *http.Client
retries int
}
type Option func(*clientConfig)
type clientConfig struct {
timeout time.Duration
retries int
rootCAs *x509.CertPool
}
func WithTimeout(timeout time.Duration) Option {
return func(cfg *clientConfig) {
cfg.timeout = timeout
}
}
func WithRetries(retries int) Option {
return func(cfg *clientConfig) {
cfg.retries = retries
}
}
func WithRootCAs(rootCAs *x509.CertPool) Option {
return func(cfg *clientConfig) {
cfg.rootCAs = rootCAs
}
}
func NewClient(cluster cluster.Cluster, opts ...Option) (*Client, error) {
baseURL, err := url.Parse(strings.TrimRight(cluster.APIEndpoint, "/"))
if err != nil {
return nil, fmt.Errorf("parse proxmox api endpoint: %w", err)
}
if baseURL.Scheme != "https" || baseURL.Host == "" {
return nil, fmt.Errorf("proxmox api endpoint must be an https url")
}
if strings.TrimSpace(cluster.TokenID) == "" {
return nil, fmt.Errorf("proxmox token id is required")
}
if strings.TrimSpace(cluster.TokenSecret) == "" {
return nil, fmt.Errorf("proxmox token secret is required")
}
fingerprint, err := normalizeFingerprint(cluster.TLSFingerprint)
if err != nil {
return nil, err
}
cfg := clientConfig{
timeout: defaultTimeout,
retries: 2,
}
for _, opt := range opts {
opt(&cfg)
}
if cfg.timeout <= 0 {
return nil, fmt.Errorf("timeout must be greater than 0")
}
if cfg.retries < 0 {
return nil, fmt.Errorf("retries must not be negative")
}
return &Client{
baseURL: baseURL,
tokenID: cluster.TokenID,
tokenSecret: cluster.TokenSecret,
httpClient: &http.Client{
Timeout: cfg.timeout,
Transport: &http.Transport{
Proxy: http.ProxyFromEnvironment,
TLSClientConfig: &tls.Config{
MinVersion: tls.VersionTLS12,
RootCAs: cfg.rootCAs,
VerifyConnection: func(state tls.ConnectionState) error {
return verifyFingerprint(state, fingerprint)
},
},
},
},
retries: cfg.retries,
}, nil
}
func (c *Client) Get(ctx context.Context, path string) (*http.Response, error) {
return c.Do(ctx, http.MethodGet, path, nil)
}
func (c *Client) Post(ctx context.Context, path string, body []byte) (*http.Response, error) {
return c.Do(ctx, http.MethodPost, path, body)
}
func (c *Client) Put(ctx context.Context, path string, body []byte) (*http.Response, error) {
return c.Do(ctx, http.MethodPut, path, body)
}
func (c *Client) Delete(ctx context.Context, path string) (*http.Response, error) {
return c.Do(ctx, http.MethodDelete, path, nil)
}
type TaskStatus struct {
Status string
ExitStatus string
}
func (c *Client) GetTaskStatus(ctx context.Context, node string, upid string) (TaskStatus, error) {
response, err := c.Get(ctx, fmt.Sprintf(
"/nodes/%s/tasks/%s/status",
url.PathEscape(node),
url.PathEscape(upid),
))
if err != nil {
return TaskStatus{}, err
}
defer response.Body.Close()
if response.StatusCode >= http.StatusBadRequest {
_, _ = io.Copy(io.Discard, response.Body)
return TaskStatus{}, fmt.Errorf("proxmox returned %s", response.Status)
}
var body struct {
Data struct {
Status string `json:"status"`
ExitStatus string `json:"exitstatus"`
} `json:"data"`
}
if err := json.NewDecoder(response.Body).Decode(&body); err != nil {
return TaskStatus{}, err
}
return TaskStatus{
Status: body.Data.Status,
ExitStatus: body.Data.ExitStatus,
}, nil
}
func (c *Client) Do(ctx context.Context, method string, path string, body []byte) (*http.Response, error) {
var lastErr error
attempts := c.retries + 1
for attempt := 0; attempt < attempts; attempt++ {
response, err := c.doOnce(ctx, method, path, body)
if err == nil && response.StatusCode < http.StatusInternalServerError {
return response, nil
}
if err == nil {
_, _ = io.Copy(io.Discard, response.Body)
_ = response.Body.Close()
lastErr = fmt.Errorf("proxmox returned %s", response.Status)
} else {
lastErr = err
}
if ctx.Err() != nil {
return nil, ctx.Err()
}
}
return nil, lastErr
}
func (c *Client) doOnce(ctx context.Context, method string, path string, body []byte) (*http.Response, error) {
requestURL := c.resolvePath(path)
request, err := http.NewRequestWithContext(ctx, method, requestURL, bytes.NewReader(body))
if err != nil {
return nil, err
}
request.Header.Set("Authorization", fmt.Sprintf("PVEAPIToken=%s=%s", c.tokenID, c.tokenSecret))
if body != nil {
request.Header.Set("Content-Type", "application/json")
}
return c.httpClient.Do(request)
}
func (c *Client) resolvePath(path string) string {
resolved := *c.baseURL
resolved.Path = joinURLPath(c.baseURL.Path, path)
return resolved.String()
}
func joinURLPath(basePath string, path string) string {
basePath = strings.TrimRight(basePath, "/")
path = strings.TrimLeft(path, "/")
if path == "" {
return basePath
}
return basePath + "/" + path
}
func normalizeFingerprint(fingerprint string) (string, error) {
normalized := strings.NewReplacer(":", "", " ", "", "-", "").Replace(strings.TrimSpace(fingerprint))
normalized = strings.ToLower(normalized)
if len(normalized) != sha256.Size*2 {
return "", fmt.Errorf("tls fingerprint must be a sha256 hex digest")
}
if _, err := hex.DecodeString(normalized); err != nil {
return "", fmt.Errorf("tls fingerprint must be hex encoded: %w", err)
}
return normalized, nil
}
func verifyFingerprint(state tls.ConnectionState, expected string) error {
if len(state.PeerCertificates) == 0 {
return fmt.Errorf("server certificate missing")
}
sum := sha256.Sum256(state.PeerCertificates[0].Raw)
actual := hex.EncodeToString(sum[:])
if subtle.ConstantTimeCompare([]byte(actual), []byte(expected)) != 1 {
return fmt.Errorf("server certificate fingerprint mismatch")
}
return nil
}
+144
View File
@@ -0,0 +1,144 @@
package proxmox
import (
"context"
"crypto/sha256"
"crypto/x509"
"encoding/hex"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"forgejo.digital-droplets.de/philschlo/proxui/platform/cluster"
)
func TestClientAcceptsMatchingFingerprintAndSendsTokenHeader(t *testing.T) {
var authHeader string
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
authHeader = r.Header.Get("Authorization")
if r.URL.Path != "/api2/json/version" {
t.Fatalf("path = %q, want /api2/json/version", r.URL.Path)
}
_, _ = w.Write([]byte(`{"data":{"version":"8.2"}}`))
}))
defer server.Close()
client := newTestClient(t, server, fingerprintForServer(server))
response, err := client.Get(context.Background(), "/version")
if err != nil {
t.Fatalf("Get() error = %v", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusOK)
}
if authHeader != "PVEAPIToken=root@pam!proxui=secret-token" {
t.Fatalf("Authorization = %q", authHeader)
}
}
func TestClientRejectsMismatchedFingerprint(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
defer server.Close()
client := newTestClient(t, server, strings.Repeat("0", sha256.Size*2))
_, err := client.Get(context.Background(), "/version")
if err == nil {
t.Fatal("Get() error = nil, want error")
}
if !strings.Contains(err.Error(), "fingerprint mismatch") {
t.Fatalf("error = %q, want fingerprint mismatch", err)
}
}
func TestClientRetriesServerErrors(t *testing.T) {
var calls int
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
if calls == 1 {
http.Error(w, "temporary", http.StatusBadGateway)
return
}
w.WriteHeader(http.StatusNoContent)
}))
defer server.Close()
client := newTestClient(t, server, fingerprintForServer(server))
response, err := client.Post(context.Background(), "/nodes/pve/status", []byte(`{"command":"start"}`))
if err != nil {
t.Fatalf("Post() error = %v", err)
}
defer response.Body.Close()
_, _ = io.Copy(io.Discard, response.Body)
if response.StatusCode != http.StatusNoContent {
t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusNoContent)
}
if calls != 2 {
t.Fatalf("calls = %d, want 2", calls)
}
}
func TestClientGetsTaskStatus(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/api2/json/nodes/pve/tasks/UPID:pve:1/status" {
t.Fatalf("path = %q", r.URL.Path)
}
_, _ = w.Write([]byte(`{"data":{"status":"stopped","exitstatus":"OK"}}`))
}))
defer server.Close()
client := newTestClient(t, server, fingerprintForServer(server))
status, err := client.GetTaskStatus(context.Background(), "pve", "UPID:pve:1")
if err != nil {
t.Fatalf("GetTaskStatus() error = %v", err)
}
if status.Status != "stopped" {
t.Fatalf("Status = %q, want stopped", status.Status)
}
if status.ExitStatus != "OK" {
t.Fatalf("ExitStatus = %q, want OK", status.ExitStatus)
}
}
func TestNewClientRejectsInvalidFingerprint(t *testing.T) {
_, err := NewClient(cluster.Cluster{
APIEndpoint: "https://pve.example.test:8006/api2/json",
TLSFingerprint: "invalid",
TokenID: "root@pam!proxui",
TokenSecret: "secret-token",
})
if err == nil {
t.Fatal("NewClient() error = nil, want error")
}
}
func newTestClient(t *testing.T, server *httptest.Server, fingerprint string) *Client {
t.Helper()
rootCAs := x509.NewCertPool()
rootCAs.AddCert(server.Certificate())
client, err := NewClient(cluster.Cluster{
APIEndpoint: server.URL + "/api2/json",
TLSFingerprint: fingerprint,
TokenID: "root@pam!proxui",
TokenSecret: "secret-token",
}, WithRootCAs(rootCAs), WithTimeout(time.Second), WithRetries(1))
if err != nil {
t.Fatalf("NewClient() error = %v", err)
}
return client
}
func fingerprintForServer(server *httptest.Server) string {
sum := sha256.Sum256(server.Certificate().Raw)
return hex.EncodeToString(sum[:])
}