Files
pxmon/internal/cluster/service_test.go
T
2026-06-16 21:52:10 +04:00

555 lines
13 KiB
Go

package cluster
import (
"context"
"os"
"os/exec"
"path/filepath"
"strings"
"testing"
"time"
)
func newTestService(t *testing.T) *Service {
t.Helper()
store, err := NewStore(filepath.Join(t.TempDir(), "clusters.enc"))
if err != nil {
t.Fatalf("NewStore error: %v", err)
}
return NewService(store)
}
func TestConnectListUseDisconnect(t *testing.T) {
t.Parallel()
svc := newTestService(t)
ctx := context.Background()
c1, _, err := svc.Connect(ctx, ConnectOptions{
Name: "eu-1",
Host: "10.0.0.10",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
Password: "pass1",
SkipCheck: true,
})
if err != nil {
t.Fatalf("connect c1 error: %v", err)
}
c2, _, err := svc.Connect(ctx, ConnectOptions{
Name: "us-1",
Host: "10.0.0.11",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
Password: "pass2",
SkipCheck: true,
})
if err != nil {
t.Fatalf("connect c2 error: %v", err)
}
clusters, activeID, err := svc.List()
if err != nil {
t.Fatalf("list error: %v", err)
}
if len(clusters) != 2 {
t.Fatalf("expected 2 clusters, got %d", len(clusters))
}
if activeID != c2.ID {
t.Fatalf("expected active %s, got %s", c2.ID, activeID)
}
_, err = svc.Use(c1.Name)
if err != nil {
t.Fatalf("use error: %v", err)
}
current, err := svc.Current()
if err != nil {
t.Fatalf("current error: %v", err)
}
if current.ID != c1.ID {
t.Fatalf("expected current %s, got %s", c1.ID, current.ID)
}
_, err = svc.Disconnect(c1.Name)
if err != nil {
t.Fatalf("disconnect error: %v", err)
}
current, err = svc.Current()
if err != nil {
t.Fatalf("current after disconnect error: %v", err)
}
if current.ID != c2.ID {
t.Fatalf("expected fallback current %s, got %s", c2.ID, current.ID)
}
}
func TestConnectDuplicateNameRequiresForce(t *testing.T) {
t.Parallel()
svc := newTestService(t)
ctx := context.Background()
_, _, err := svc.Connect(ctx, ConnectOptions{
Name: "prod",
Host: "10.0.0.10",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
Password: "pass1",
SkipCheck: true,
})
if err != nil {
t.Fatalf("initial connect error: %v", err)
}
_, _, err = svc.Connect(ctx, ConnectOptions{
Name: "prod",
Host: "10.0.0.20",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
Password: "pass2",
SkipCheck: true,
})
if err == nil {
t.Fatal("expected duplicate name error")
}
c, _, err := svc.Connect(ctx, ConnectOptions{
Name: "prod",
Host: "10.0.0.20",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
Password: "pass2",
SkipCheck: true,
Force: true,
})
if err != nil {
t.Fatalf("force connect error: %v", err)
}
if c.Host != "10.0.0.20" {
t.Fatalf("expected overwritten host, got %s", c.Host)
}
}
func TestExportImportRestoresKeyFilesAndNewFields(t *testing.T) {
t.Parallel()
srcDir := t.TempDir()
keyPath := filepath.Join(srcDir, "id_ed25519")
passPath := filepath.Join(srcDir, "pass.pxmonpassphrase")
sftpKeyPath := filepath.Join(srcDir, "sftp_key")
if err := os.WriteFile(keyPath, []byte("PRIVATE KEY\n"), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(passPath, []byte("secret-pass\n"), 0o600); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(sftpKeyPath, []byte("SFTP PRIVATE KEY\n"), 0o600); err != nil {
t.Fatal(err)
}
src := newTestService(t)
reg := newRegistry()
reg.ActiveClusterID = "clu_1"
reg.Backups.Targets = []BackupTarget{{
ID: "bt_1",
Name: "sftp",
Type: "sftp",
Enabled: true,
SFTPKeyPath: sftpKeyPath,
}}
reg.Clusters = []Cluster{{
ID: "clu_1",
Name: "node",
Host: "192.0.2.10",
Port: 22,
User: "root",
Transport: TransportIPFabric,
AuthMethod: AuthMethodKey,
KeyPath: keyPath,
KeyPassphraseFile: passPath,
RepoTunnel: RepoTunnelState{
Enabled: true,
Proxy: "http://203.0.113.10:3128",
Source: "dnf",
},
Alerts: defaultAlertPolicy(),
CreatedAt: time.Now().UTC(),
UpdatedAt: time.Now().UTC(),
}}
if err := src.store.Save(reg); err != nil {
t.Fatalf("save source registry: %v", err)
}
bundlePath := filepath.Join(t.TempDir(), "pxmon-export.enc")
if err := src.Export(bundlePath, "strong-test-passphrase"); err != nil {
t.Fatalf("export error: %v", err)
}
raw, err := os.ReadFile(bundlePath)
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(raw), "PRIVATE KEY") || strings.Contains(string(raw), "secret-pass") {
t.Fatal("export bundle leaked key material in plaintext")
}
dst := newTestService(t)
if _, err := dst.Import(bundlePath, "strong-test-passphrase", ImportModeReplace); err != nil {
t.Fatalf("import error: %v", err)
}
gotReg, err := dst.store.Load()
if err != nil {
t.Fatalf("load imported registry: %v", err)
}
if len(gotReg.Clusters) != 1 {
t.Fatalf("expected 1 cluster, got %d", len(gotReg.Clusters))
}
got := gotReg.Clusters[0]
if !got.RepoTunnel.Enabled || got.RepoTunnel.Proxy != "http://203.0.113.10:3128" {
t.Fatalf("repo tunnel was not preserved: %+v", got.RepoTunnel)
}
if got.KeyPath == keyPath || got.KeyPassphraseFile == passPath {
t.Fatalf("expected key paths to be restored under destination config dir, got key=%q pass=%q", got.KeyPath, got.KeyPassphraseFile)
}
keyData, err := os.ReadFile(got.KeyPath)
if err != nil {
t.Fatalf("read restored key: %v", err)
}
if string(keyData) != "PRIVATE KEY\n" {
t.Fatalf("unexpected restored key content: %q", string(keyData))
}
passData, err := os.ReadFile(got.KeyPassphraseFile)
if err != nil {
t.Fatalf("read restored passphrase file: %v", err)
}
if string(passData) != "secret-pass\n" {
t.Fatalf("unexpected restored passphrase content: %q", string(passData))
}
if len(gotReg.Backups.Targets) != 1 || gotReg.Backups.Targets[0].SFTPKeyPath == sftpKeyPath {
t.Fatalf("expected restored SFTP key path, got %+v", gotReg.Backups.Targets)
}
mergeDst := newTestService(t)
if _, err := mergeDst.Import(bundlePath, "strong-test-passphrase", ImportModeMerge); err != nil {
t.Fatalf("merge import error: %v", err)
}
mergeReg, err := mergeDst.store.Load()
if err != nil {
t.Fatalf("load merge registry: %v", err)
}
if len(mergeReg.Backups.Targets) != 1 || strings.TrimSpace(mergeReg.Backups.Targets[0].SFTPKeyPath) == "" {
t.Fatalf("expected backup target to be merged, got %+v", mergeReg.Backups)
}
}
func TestConnectRequiresAuthData(t *testing.T) {
t.Parallel()
svc := newTestService(t)
ctx := context.Background()
_, _, err := svc.Connect(ctx, ConnectOptions{
Name: "bad",
Host: "10.0.0.10",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
SkipCheck: true,
})
if err == nil {
t.Fatal("expected error for empty password auth")
}
}
func TestAlertPolicySetAndGet(t *testing.T) {
t.Parallel()
svc := newTestService(t)
ctx := context.Background()
_, _, err := svc.Connect(ctx, ConnectOptions{
Name: "node-1",
Host: "10.0.0.10",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
Password: "pass",
SkipCheck: true,
})
if err != nil {
t.Fatalf("connect error: %v", err)
}
updated, err := svc.SetAlertPolicy("node-1", AlertPolicy{
CPUWarnPercent: 70,
RAMWarnPercent: 75,
SwapWarnPercent: 60,
DiskWarnPercent: 80,
NetWarnMbps: 120,
NetSustainEnabled: true,
NetSustainIface: "eth0",
NetSustainInclude: []string{"net0"},
NetSustainExclude: []string{"backup"},
NetSustainMbps: 500,
NetSustainMinutes: 60,
NetSustainCooldownMins: 15,
})
if err != nil {
t.Fatalf("set alert policy error: %v", err)
}
if updated.Alerts.NetWarnMbps != 120 {
t.Fatalf("unexpected net threshold: %.2f", updated.Alerts.NetWarnMbps)
}
got, err := svc.GetAlertPolicy("node-1")
if err != nil {
t.Fatalf("get alert policy error: %v", err)
}
if got.RAMWarnPercent != 75 || got.DiskWarnPercent != 80 {
t.Fatalf("unexpected thresholds: %+v", got)
}
if !got.NetSustainEnabled || got.NetSustainIface != "eth0" || got.NetSustainMbps != 500 {
t.Fatalf("unexpected sustained net policy: %+v", got)
}
if len(got.NetSustainInclude) != 1 || got.NetSustainInclude[0] != "net0" {
t.Fatalf("unexpected sustained net include filter: %+v", got.NetSustainInclude)
}
if len(got.NetSustainExclude) != 1 || got.NetSustainExclude[0] != "backup" {
t.Fatalf("unexpected sustained net exclude filter: %+v", got.NetSustainExclude)
}
}
func TestIsPluginToolSupported(t *testing.T) {
t.Parallel()
info := SoftwareInfo{
Bird: true,
FRR: true,
KVM: true,
LXC: true,
LXD: false,
}
tests := []struct {
tool string
ok bool
}{
{tool: "bird", ok: true},
{tool: "frr", ok: true},
{tool: "kvm", ok: true},
{tool: "lxc", ok: true},
{tool: "lxd", ok: true}, // lxd aliases to lxc command templates
{tool: "unknown", ok: false},
}
for _, tt := range tests {
tt := tt
t.Run(tt.tool, func(t *testing.T) {
t.Parallel()
if got := isPluginToolSupported(info, tt.tool); got != tt.ok {
t.Fatalf("tool=%s expected %v got %v", tt.tool, tt.ok, got)
}
})
}
}
func TestLXDTopUsesTemplateNotRawLxcTop(t *testing.T) {
t.Parallel()
script, err := pluginScript("lxd", "top", nil)
if err != nil {
t.Fatalf("pluginScript error: %v", err)
}
if strings.Contains(script, "lxc top") {
t.Fatalf("expected custom template script, got raw lxc top: %q", script)
}
if !strings.Contains(script, "lxc info") {
t.Fatalf("expected lxc info usage in top template")
}
}
func TestRunPluginActionReportsMissingSupport(t *testing.T) {
t.Parallel()
svc := newTestService(t)
ctx := context.Background()
cluster, _, err := svc.Connect(ctx, ConnectOptions{
Name: "eu-1",
Host: "127.0.0.1",
Port: 22,
User: "root",
AuthMethod: AuthMethodPassword,
Password: "pass",
SkipCheck: true,
})
if err != nil {
t.Fatalf("connect error: %v", err)
}
reg, err := svc.store.Load()
if err != nil {
t.Fatalf("load registry error: %v", err)
}
for i := range reg.Clusters {
if reg.Clusters[i].ID == cluster.ID {
reg.Clusters[i].Software = SoftwareInfo{
DetectedAt: time.Now().UTC(),
}
break
}
}
if err := svc.store.Save(reg); err != nil {
t.Fatalf("save registry error: %v", err)
}
_, err = svc.RunPluginAction(ctx, cluster.Name, "lxd", "top", nil)
if err == nil {
t.Fatalf("expected unsupported software error")
}
if !strings.Contains(err.Error(), "support for lxd was not detected") {
t.Fatalf("unexpected error message: %v", err)
}
}
func TestKVMTopScriptIncludesReadableMetrics(t *testing.T) {
t.Parallel()
script, err := pluginScript("kvm", "top", nil)
if err != nil {
t.Fatalf("pluginScript error: %v", err)
}
for _, want := range []string{
"VCPU",
"RAM_MAX",
"DISK_CAP",
"DISK_ALLOC",
} {
if !strings.Contains(script, want) {
t.Fatalf("expected %q in kvm top script", want)
}
}
}
func TestKVMNetTopScriptIncludesRateAndP95(t *testing.T) {
t.Parallel()
script, err := pluginScript("kvm", "net-top", nil)
if err != nil {
t.Fatalf("pluginScript error: %v", err)
}
for _, want := range []string{
"NET_Mbps",
"RX_Mbps",
"TX_Mbps",
"P95_Mbps",
"RX_TOTAL",
"TX_TOTAL",
} {
if !strings.Contains(script, want) {
t.Fatalf("expected %q in kvm net-top script", want)
}
}
}
func TestKVMScriptsAreShellParseable(t *testing.T) {
t.Parallel()
cases := []struct {
tool string
action string
}{
{tool: "kvm", action: "top"},
{tool: "kvm", action: "net-top"},
}
for _, tc := range cases {
tc := tc
t.Run(tc.tool+"-"+tc.action, func(t *testing.T) {
t.Parallel()
script, err := pluginScript(tc.tool, tc.action, nil)
if err != nil {
t.Fatalf("pluginScript error: %v", err)
}
cmd := exec.Command("sh", "-n", "-c", script)
out, err := cmd.CombinedOutput()
if err != nil {
t.Fatalf("shell parse failed: %v\n%s\nSCRIPT:\n%s", err, string(out), script)
}
})
}
}
func TestTelegramConfigSetGetDisable(t *testing.T) {
t.Parallel()
svc := newTestService(t)
cfg, err := svc.SetTelegram(Telegram{
Enabled: true,
Token: "123:ABC",
AllowedUserIDs: []int64{2002, 1001, 2002},
})
if err != nil {
t.Fatalf("set telegram config: %v", err)
}
if !cfg.Enabled {
t.Fatal("expected enabled telegram config")
}
if len(cfg.AllowedUserIDs) != 2 {
t.Fatalf("expected deduped ids, got %+v", cfg.AllowedUserIDs)
}
got, err := svc.GetTelegram()
if err != nil {
t.Fatalf("get telegram config: %v", err)
}
if got.Token != "123:ABC" {
t.Fatalf("unexpected token: %q", got.Token)
}
if len(got.AllowedUserIDs) != 2 || got.AllowedUserIDs[0] != 1001 || got.AllowedUserIDs[1] != 2002 {
t.Fatalf("unexpected allowed ids: %+v", got.AllowedUserIDs)
}
disabled, err := svc.DisableTelegram()
if err != nil {
t.Fatalf("disable telegram config: %v", err)
}
if disabled.Enabled {
t.Fatal("expected disabled telegram config")
}
}
func TestTelegramConfigValidation(t *testing.T) {
t.Parallel()
svc := newTestService(t)
_, err := svc.SetTelegram(Telegram{
Enabled: true,
AllowedUserIDs: []int64{123},
})
if err == nil {
t.Fatal("expected token validation error")
}
_, err = svc.SetTelegram(Telegram{
Enabled: true,
Token: "123:ABC",
})
if err == nil {
t.Fatal("expected allowed ids validation error")
}
}