feat: add proxmox task polling job

This commit is contained in:
Philipp
2026-06-11 10:07:17 +02:00
parent e34dc28fd3
commit 7c818cf92e
6 changed files with 488 additions and 2 deletions
+11 -2
View File
@@ -43,7 +43,7 @@ func main() {
Concurrency: cfg.WorkerConcurrency,
})
mux := asynq.NewServeMux()
registerHandlers(mux, logger)
registerHandlers(mux, db, logger)
go func() {
logger.Info("worker started", "concurrency", cfg.WorkerConcurrency, "redis_addr", cfg.RedisAddr)
@@ -58,8 +58,17 @@ func main() {
logger.Info("worker stopped")
}
func registerHandlers(mux *asynq.ServeMux, logger *slog.Logger) {
func registerHandlers(mux *asynq.ServeMux, db *sql.DB, logger *slog.Logger) {
mux.Handle(tasks.TypeDummy, tasks.NewDummyHandler(logger))
mux.Handle(
tasks.TypeProxmoxTaskPoll,
tasks.NewProxmoxTaskPollHandler(
tasks.UnconfiguredProxmoxTaskClientResolver{},
tasks.NewSQLVMStatusStore(db),
tasks.NewSQLAuditWriter(db),
logger,
),
)
}
func openDatabase(databaseURL string) (*sql.DB, error) {
+285
View File
@@ -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")
}
+183
View File
@@ -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")
}
}