feat: add proxmox task polling job
This commit is contained in:
@@ -0,0 +1,285 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
const TypeProxmoxTaskPoll = "proxmox.task.poll"
|
||||
|
||||
type ProxmoxTaskPollPayload struct {
|
||||
ClusterID string `json:"cluster_id"`
|
||||
Node string `json:"node"`
|
||||
UPID string `json:"upid"`
|
||||
TargetType string `json:"target_type"`
|
||||
TargetID string `json:"target_id"`
|
||||
TenantID string `json:"tenant_id"`
|
||||
ProfileID string `json:"profile_id,omitempty"`
|
||||
Action string `json:"action"`
|
||||
SuccessStatus string `json:"success_status"`
|
||||
}
|
||||
|
||||
type ProxmoxTaskStatus struct {
|
||||
Status string
|
||||
ExitStatus string
|
||||
}
|
||||
|
||||
type ProxmoxTaskClient interface {
|
||||
GetTaskStatus(ctx context.Context, node string, upid string) (ProxmoxTaskStatus, error)
|
||||
}
|
||||
|
||||
type ProxmoxTaskClientResolver interface {
|
||||
ResolveTaskClient(ctx context.Context, clusterID string) (ProxmoxTaskClient, error)
|
||||
}
|
||||
|
||||
type VMStatusStore interface {
|
||||
SetVMStatus(ctx context.Context, vmID string, status string) error
|
||||
}
|
||||
|
||||
type AuditWriter interface {
|
||||
WriteAudit(ctx context.Context, event AuditEvent) error
|
||||
}
|
||||
|
||||
type AuditEvent struct {
|
||||
TenantID string
|
||||
ProfileID string
|
||||
Action string
|
||||
TargetType string
|
||||
TargetID string
|
||||
Metadata map[string]any
|
||||
}
|
||||
|
||||
type ProxmoxTaskPollHandler struct {
|
||||
resolver ProxmoxTaskClientResolver
|
||||
vmStore VMStatusStore
|
||||
auditWriter AuditWriter
|
||||
logger *slog.Logger
|
||||
maxPolls int
|
||||
pollDelay time.Duration
|
||||
}
|
||||
|
||||
func NewProxmoxTaskPollHandler(
|
||||
resolver ProxmoxTaskClientResolver,
|
||||
vmStore VMStatusStore,
|
||||
auditWriter AuditWriter,
|
||||
logger *slog.Logger,
|
||||
) ProxmoxTaskPollHandler {
|
||||
return ProxmoxTaskPollHandler{
|
||||
resolver: resolver,
|
||||
vmStore: vmStore,
|
||||
auditWriter: auditWriter,
|
||||
logger: logger,
|
||||
maxPolls: 60,
|
||||
pollDelay: 2 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
func (h ProxmoxTaskPollHandler) WithPolling(maxPolls int, pollDelay time.Duration) ProxmoxTaskPollHandler {
|
||||
h.maxPolls = maxPolls
|
||||
h.pollDelay = pollDelay
|
||||
return h
|
||||
}
|
||||
|
||||
func NewProxmoxTaskPollTask(payload ProxmoxTaskPollPayload) (*asynq.Task, error) {
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return asynq.NewTask(TypeProxmoxTaskPoll, body), nil
|
||||
}
|
||||
|
||||
func (h ProxmoxTaskPollHandler) ProcessTask(ctx context.Context, task *asynq.Task) error {
|
||||
payload, err := decodeProxmoxTaskPollPayload(task.Payload())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
client, err := h.resolver.ResolveTaskClient(ctx, payload.ClusterID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve proxmox task client: %w", err)
|
||||
}
|
||||
|
||||
for attempt := 0; attempt < h.maxPolls; attempt++ {
|
||||
status, err := client.GetTaskStatus(ctx, payload.Node, payload.UPID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("poll proxmox task status: %w", err)
|
||||
}
|
||||
|
||||
if isTaskRunning(status) {
|
||||
if err := sleepContext(ctx, h.pollDelay); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if isTaskSuccessful(status) {
|
||||
return h.completeTarget(ctx, payload, payload.SuccessStatus, status)
|
||||
}
|
||||
|
||||
return h.completeTarget(ctx, payload, "failed", status)
|
||||
}
|
||||
|
||||
return fmt.Errorf("proxmox task %s still running after %d polls", payload.UPID, h.maxPolls)
|
||||
}
|
||||
|
||||
func (h ProxmoxTaskPollHandler) completeTarget(ctx context.Context, payload ProxmoxTaskPollPayload, finalStatus string, taskStatus ProxmoxTaskStatus) error {
|
||||
if payload.TargetType != "vm" {
|
||||
return fmt.Errorf("unsupported target type %q", payload.TargetType)
|
||||
}
|
||||
|
||||
if err := h.vmStore.SetVMStatus(ctx, payload.TargetID, finalStatus); err != nil {
|
||||
return fmt.Errorf("set vm status: %w", err)
|
||||
}
|
||||
|
||||
if err := h.auditWriter.WriteAudit(ctx, AuditEvent{
|
||||
TenantID: payload.TenantID,
|
||||
ProfileID: payload.ProfileID,
|
||||
Action: payload.Action,
|
||||
TargetType: payload.TargetType,
|
||||
TargetID: payload.TargetID,
|
||||
Metadata: map[string]any{
|
||||
"cluster_id": payload.ClusterID,
|
||||
"node": payload.Node,
|
||||
"upid": payload.UPID,
|
||||
"task_status": taskStatus.Status,
|
||||
"exit_status": taskStatus.ExitStatus,
|
||||
"target_status": finalStatus,
|
||||
"success_status": payload.SuccessStatus,
|
||||
},
|
||||
}); err != nil {
|
||||
return fmt.Errorf("write audit event: %w", err)
|
||||
}
|
||||
|
||||
h.logger.Info("completed proxmox task poll", "target_type", payload.TargetType, "target_id", payload.TargetID, "status", finalStatus)
|
||||
return nil
|
||||
}
|
||||
|
||||
func decodeProxmoxTaskPollPayload(body []byte) (ProxmoxTaskPollPayload, error) {
|
||||
var payload ProxmoxTaskPollPayload
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("decode proxmox task poll payload: %w", err)
|
||||
}
|
||||
if payload.ClusterID == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("cluster_id is required")
|
||||
}
|
||||
if payload.Node == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("node is required")
|
||||
}
|
||||
if payload.UPID == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("upid is required")
|
||||
}
|
||||
if payload.TargetType == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("target_type is required")
|
||||
}
|
||||
if payload.TargetID == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("target_id is required")
|
||||
}
|
||||
if payload.TenantID == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("tenant_id is required")
|
||||
}
|
||||
if payload.Action == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("action is required")
|
||||
}
|
||||
if payload.SuccessStatus == "" {
|
||||
return ProxmoxTaskPollPayload{}, fmt.Errorf("success_status is required")
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func isTaskRunning(status ProxmoxTaskStatus) bool {
|
||||
return status.Status == "" || status.Status == "running"
|
||||
}
|
||||
|
||||
func isTaskSuccessful(status ProxmoxTaskStatus) bool {
|
||||
return status.Status == "stopped" && status.ExitStatus == "OK"
|
||||
}
|
||||
|
||||
func sleepContext(ctx context.Context, delay time.Duration) error {
|
||||
if delay <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
timer := time.NewTimer(delay)
|
||||
defer timer.Stop()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-timer.C:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
type SQLVMStatusStore struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewSQLVMStatusStore(db *sql.DB) SQLVMStatusStore {
|
||||
return SQLVMStatusStore{db: db}
|
||||
}
|
||||
|
||||
func (s SQLVMStatusStore) SetVMStatus(ctx context.Context, vmID string, status string) error {
|
||||
result, err := s.db.ExecContext(ctx, `
|
||||
update public.vms
|
||||
set status = $2,
|
||||
updated_at = now()
|
||||
where id = $1
|
||||
`, vmID, status)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rows == 0 {
|
||||
return fmt.Errorf("vm not found")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type SQLAuditWriter struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewSQLAuditWriter(db *sql.DB) SQLAuditWriter {
|
||||
return SQLAuditWriter{db: db}
|
||||
}
|
||||
|
||||
func (w SQLAuditWriter) WriteAudit(ctx context.Context, event AuditEvent) error {
|
||||
metadata, err := json.Marshal(event.Metadata)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var profileID any
|
||||
if event.ProfileID != "" {
|
||||
profileID = event.ProfileID
|
||||
}
|
||||
|
||||
_, err = w.db.ExecContext(ctx, `
|
||||
insert into public.audit_log (
|
||||
tenant_id,
|
||||
profile_id,
|
||||
action,
|
||||
target_type,
|
||||
target_id,
|
||||
metadata
|
||||
)
|
||||
values ($1, $2, $3, $4, $5, $6::jsonb)
|
||||
`, event.TenantID, profileID, event.Action, event.TargetType, event.TargetID, string(metadata))
|
||||
return err
|
||||
}
|
||||
|
||||
type UnconfiguredProxmoxTaskClientResolver struct{}
|
||||
|
||||
func (UnconfiguredProxmoxTaskClientResolver) ResolveTaskClient(context.Context, string) (ProxmoxTaskClient, error) {
|
||||
return nil, fmt.Errorf("proxmox task client resolver is not configured")
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package tasks
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/hibiken/asynq"
|
||||
)
|
||||
|
||||
func TestProxmoxTaskPollHandlerMarksVMRunningOnOK(t *testing.T) {
|
||||
resolver := &stubTaskClientResolver{client: &stubTaskClient{
|
||||
statuses: []ProxmoxTaskStatus{
|
||||
{Status: "running"},
|
||||
{Status: "stopped", ExitStatus: "OK"},
|
||||
},
|
||||
}}
|
||||
vmStore := &stubVMStatusStore{}
|
||||
auditWriter := &stubAuditWriter{}
|
||||
handler := NewProxmoxTaskPollHandler(resolver, vmStore, auditWriter, slog.Default()).WithPolling(3, 0)
|
||||
|
||||
task := newPollTask(t, ProxmoxTaskPollPayload{
|
||||
ClusterID: "cluster-1",
|
||||
Node: "pve",
|
||||
UPID: "UPID:pve:1",
|
||||
TargetType: "vm",
|
||||
TargetID: "vm-1",
|
||||
TenantID: "tenant-1",
|
||||
ProfileID: "profile-1",
|
||||
Action: "vm.power.start",
|
||||
SuccessStatus: "running",
|
||||
})
|
||||
|
||||
if err := handler.ProcessTask(context.Background(), task); err != nil {
|
||||
t.Fatalf("ProcessTask() error = %v", err)
|
||||
}
|
||||
|
||||
if resolver.clusterID != "cluster-1" {
|
||||
t.Fatalf("clusterID = %q, want cluster-1", resolver.clusterID)
|
||||
}
|
||||
if vmStore.status != "running" {
|
||||
t.Fatalf("vm status = %q, want running", vmStore.status)
|
||||
}
|
||||
if auditWriter.event.Action != "vm.power.start" {
|
||||
t.Fatalf("audit action = %q, want vm.power.start", auditWriter.event.Action)
|
||||
}
|
||||
if auditWriter.event.Metadata["exit_status"] != "OK" {
|
||||
t.Fatalf("audit exit_status = %v, want OK", auditWriter.event.Metadata["exit_status"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxmoxTaskPollHandlerMarksVMFailedOnErrorExit(t *testing.T) {
|
||||
resolver := &stubTaskClientResolver{client: &stubTaskClient{
|
||||
statuses: []ProxmoxTaskStatus{
|
||||
{Status: "stopped", ExitStatus: "error"},
|
||||
},
|
||||
}}
|
||||
vmStore := &stubVMStatusStore{}
|
||||
auditWriter := &stubAuditWriter{}
|
||||
handler := NewProxmoxTaskPollHandler(resolver, vmStore, auditWriter, slog.Default()).WithPolling(2, 0)
|
||||
|
||||
task := newPollTask(t, validPollPayload())
|
||||
|
||||
if err := handler.ProcessTask(context.Background(), task); err != nil {
|
||||
t.Fatalf("ProcessTask() error = %v", err)
|
||||
}
|
||||
|
||||
if vmStore.status != "failed" {
|
||||
t.Fatalf("vm status = %q, want failed", vmStore.status)
|
||||
}
|
||||
if auditWriter.event.Metadata["target_status"] != "failed" {
|
||||
t.Fatalf("audit target_status = %v, want failed", auditWriter.event.Metadata["target_status"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxmoxTaskPollHandlerReturnsErrorWhenStillRunning(t *testing.T) {
|
||||
resolver := &stubTaskClientResolver{client: &stubTaskClient{
|
||||
statuses: []ProxmoxTaskStatus{
|
||||
{Status: "running"},
|
||||
{Status: "running"},
|
||||
},
|
||||
}}
|
||||
handler := NewProxmoxTaskPollHandler(resolver, &stubVMStatusStore{}, &stubAuditWriter{}, slog.Default()).WithPolling(2, 0)
|
||||
|
||||
err := handler.ProcessTask(context.Background(), newPollTask(t, validPollPayload()))
|
||||
if err == nil {
|
||||
t.Fatal("ProcessTask() error = nil, want error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxmoxTaskPollHandlerRejectsInvalidPayload(t *testing.T) {
|
||||
handler := NewProxmoxTaskPollHandler(&stubTaskClientResolver{}, &stubVMStatusStore{}, &stubAuditWriter{}, slog.Default())
|
||||
|
||||
err := handler.ProcessTask(context.Background(), asynq.NewTask(TypeProxmoxTaskPoll, []byte(`{"cluster_id":""}`)))
|
||||
if err == nil {
|
||||
t.Fatal("ProcessTask() error = nil, want error")
|
||||
}
|
||||
}
|
||||
|
||||
func newPollTask(t *testing.T, payload ProxmoxTaskPollPayload) *asynq.Task {
|
||||
t.Helper()
|
||||
|
||||
task, err := NewProxmoxTaskPollTask(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("NewProxmoxTaskPollTask() error = %v", err)
|
||||
}
|
||||
return task
|
||||
}
|
||||
|
||||
func validPollPayload() ProxmoxTaskPollPayload {
|
||||
return ProxmoxTaskPollPayload{
|
||||
ClusterID: "cluster-1",
|
||||
Node: "pve",
|
||||
UPID: "UPID:pve:1",
|
||||
TargetType: "vm",
|
||||
TargetID: "vm-1",
|
||||
TenantID: "tenant-1",
|
||||
Action: "vm.power.reboot",
|
||||
SuccessStatus: "running",
|
||||
}
|
||||
}
|
||||
|
||||
type stubTaskClientResolver struct {
|
||||
clusterID string
|
||||
client ProxmoxTaskClient
|
||||
}
|
||||
|
||||
func (s *stubTaskClientResolver) ResolveTaskClient(_ context.Context, clusterID string) (ProxmoxTaskClient, error) {
|
||||
s.clusterID = clusterID
|
||||
return s.client, nil
|
||||
}
|
||||
|
||||
type stubTaskClient struct {
|
||||
statuses []ProxmoxTaskStatus
|
||||
calls int
|
||||
}
|
||||
|
||||
func (s *stubTaskClient) GetTaskStatus(_ context.Context, node string, upid string) (ProxmoxTaskStatus, error) {
|
||||
s.calls++
|
||||
if node != "pve" {
|
||||
return ProxmoxTaskStatus{}, nil
|
||||
}
|
||||
if upid == "" {
|
||||
return ProxmoxTaskStatus{}, nil
|
||||
}
|
||||
if len(s.statuses) == 0 {
|
||||
return ProxmoxTaskStatus{Status: "running"}, nil
|
||||
}
|
||||
|
||||
status := s.statuses[0]
|
||||
s.statuses = s.statuses[1:]
|
||||
return status, nil
|
||||
}
|
||||
|
||||
type stubVMStatusStore struct {
|
||||
vmID string
|
||||
status string
|
||||
}
|
||||
|
||||
func (s *stubVMStatusStore) SetVMStatus(_ context.Context, vmID string, status string) error {
|
||||
s.vmID = vmID
|
||||
s.status = status
|
||||
return nil
|
||||
}
|
||||
|
||||
type stubAuditWriter struct {
|
||||
event AuditEvent
|
||||
}
|
||||
|
||||
func (s *stubAuditWriter) WriteAudit(_ context.Context, event AuditEvent) error {
|
||||
s.event = event
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestSleepContextReturnsContextError(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
if err := sleepContext(ctx, time.Second); err == nil {
|
||||
t.Fatal("sleepContext() error = nil, want context error")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user