Files

193 lines
5.5 KiB
Go

package nodecontrol
import (
"context"
"encoding/json"
"errors"
"io"
"net"
"net/http"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
func validRequest() Request {
return Request{Action: "join", Name: "worker-01", SSHTarget: "[email protected]", ExternalIP: "192.0.2.10", ConfirmMaintenance: true}
}
func TestRequestRejectsUnsafeInputs(t *testing.T) {
for _, change := range []func(*Request){func(r *Request) { r.ConfirmMaintenance = false }, func(r *Request) { r.SSHTarget = "-oProxyCommand=sh" }, func(r *Request) { r.SSHTarget = "root@host;id" }, func(r *Request) { r.Name = "../../escape" }, func(r *Request) { r.ExternalIP = "127.0.0.1" }, func(r *Request) { r.Peers = []string{"10.0.0.0/8"} }, func(r *Request) { r.Action = "exec" }} {
r := validRequest()
change(&r)
if r.Validate() == nil {
t.Fatalf("accepted %+v", r)
}
}
if err := validRequest().Validate(); err != nil {
t.Fatal(err)
}
}
func TestNodeNames(t *testing.T) {
for _, name := range []string{"worker-01", "worker.localdomain"} {
r := validRequest()
r.Name = name
if err := r.Validate(); err != nil {
t.Fatalf("%s: %v", name, err)
}
}
for _, name := range []string{"worker..localdomain", "worker.-domain", "worker.", strings.Repeat("a", 64)} {
r := validRequest()
r.Name = name
if r.Validate() == nil {
t.Fatalf("accepted %s", name)
}
}
}
func TestPersistentSerializedTasksAndLogs(t *testing.T) {
dir := t.TempDir()
entered := make(chan struct{})
finish := make(chan struct{})
manager, err := Open(context.Background(), dir, func(ctx context.Context, r Request, stage func(string) error, out io.Writer) error {
if err := stage("approval"); err != nil {
return err
}
io.WriteString(out, "probe failed\n")
close(entered)
<-finish
return errors.New("quarantine retained")
})
if err != nil {
t.Fatal(err)
}
task, err := manager.Start(validRequest(), "owner")
if err != nil {
t.Fatal(err)
}
<-entered
if _, err = manager.Start(validRequest(), "owner"); !errors.Is(err, ErrBusy) {
t.Fatal("concurrent task", err)
}
live, err := manager.Get(task.ID)
if err != nil || live.Stage != "approval" || !strings.Contains(live.Log, "probe failed") {
t.Fatal(live, err)
}
close(finish)
manager.Wait()
reopened, err := Open(context.Background(), dir, nil)
if err != nil {
t.Fatal(err)
}
got, err := reopened.Get(task.ID)
if err != nil || got.State != "failed" || got.Error != "quarantine retained" || got.FinishedAt == nil || got.Actor != "owner" {
t.Fatal(got, err)
}
if reopened.List()[0].Log != "" {
t.Fatal("task list leaked logs")
}
}
func TestInterruptedTaskIsNotReportedRunning(t *testing.T) {
dir := t.TempDir()
task := Task{ID: strings.Repeat("a", 32), Request: validRequest(), State: "running", StartedAt: time.Now()}
raw, _ := json.Marshal(task)
if err := os.WriteFile(filepath.Join(dir, task.ID+".json"), raw, 0600); err != nil {
t.Fatal(err)
}
manager, err := Open(context.Background(), dir, nil)
if err != nil {
t.Fatal(err)
}
got, err := manager.Get(task.ID)
if err != nil || got.State != "failed" || got.FinishedAt == nil || !strings.Contains(got.Error, "restarted") {
t.Fatal(got, err)
}
}
func TestLocalClientAndBoundedOutput(t *testing.T) {
dir := t.TempDir()
manager, err := Open(context.Background(), dir, func(ctx context.Context, r Request, stage func(string) error, out io.Writer) error {
for range 90 {
if _, err := io.WriteString(out, strings.Repeat("x", MaxLog)); err != nil {
return err
}
}
io.WriteString(out, "last output")
return nil
})
if err != nil {
t.Fatal(err)
}
// Avoid macOS' short Unix socket path limit under t.TempDir.
sock, err := os.CreateTemp("/tmp", "felis-node-test-")
if err != nil {
t.Fatal(err)
}
path := sock.Name()
sock.Close()
os.Remove(path)
defer os.Remove(path)
listener, err := net.Listen("unix", path)
if err != nil {
t.Fatal(err)
}
server := &http.Server{Handler: manager.Handler()}
go server.Serve(listener)
defer server.Close()
client := NewClient(path)
task, err := client.Start(context.Background(), validRequest(), "owner")
if err != nil {
t.Fatal(err)
}
manager.Wait()
got, err := client.Get(context.Background(), task.ID)
if err != nil || got.State != "succeeded" || len(got.Log) > MaxLog || !strings.HasSuffix(got.Log, "last output") {
t.Fatal("task/log", got.State, len(got.Log), err)
}
stat, err := os.Stat(filepath.Join(dir, task.ID+".log"))
if err != nil || stat.Size() > 4<<20 {
t.Fatal("unbounded file", stat, err)
}
if _, err := client.Get(context.Background(), "../secret"); !errors.Is(err, ErrNotFound) {
t.Fatal(err)
}
}
func TestCancellationFinishesTask(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
entered := make(chan struct{})
manager, err := Open(ctx, t.TempDir(), func(ctx context.Context, r Request, stage func(string) error, out io.Writer) error {
close(entered)
<-ctx.Done()
return ctx.Err()
})
if err != nil {
t.Fatal(err)
}
task, err := manager.Start(validRequest(), "owner")
if err != nil {
t.Fatal(err)
}
<-entered
cancel()
manager.Wait()
got, err := manager.Get(task.ID)
if err != nil || got.State != "failed" {
t.Fatal(got, err)
}
}
// Optional live check under the API container's UID and SELinux domain.
func TestHostSocketConnectivity(t *testing.T) {
socket := os.Getenv("FELIS_TEST_NODE_CONTROL_SOCKET")
if socket == "" {
t.Skip("live host socket not supplied")
}
tasks, err := NewClient(socket).List(context.Background())
if err != nil {
t.Fatal(err)
}
if tasks == nil {
t.Fatal("host returned no task list")
}
}