From cad3b6b01a1a7f03dea86f64bdb835fbfe2f81cd Mon Sep 17 00:00:00 2001 From: IOHelpMe <101821466+IOHelpMe@users.noreply.github.com> Date: Fri, 24 Jul 2026 01:59:46 +0200 Subject: [PATCH 1/4] feat(runtime): expose scroll pull SnapshotProgress --- apps/druid/adapters/cli/callback.go | 17 ++++ apps/druid/adapters/cli/callback_test.go | 40 +++++++++ apps/druid/adapters/cli/daemon.go | 6 +- .../adapters/cli/worker_progress_test.go | 53 ++++++++++++ apps/druid/adapters/cli/worker_pull.go | 81 ++++++++++++++++--- apps/druid/adapters/cli/worker_test.go | 2 +- .../adapters/http/handlers/health_handler.go | 19 ++++- .../http/handlers/health_handler_test.go | 33 ++++++++ apps/druid/core/services/worker_callbacks.go | 34 +++++++- .../core/services/worker_callbacks_test.go | 26 ++++++ internal/core/services/registry/oci.go | 29 +++---- internal/core/services/registry/oci_test.go | 41 ++++++++++ .../core/services/runtime_scroll_manager.go | 6 +- 13 files changed, 353 insertions(+), 34 deletions(-) create mode 100644 apps/druid/adapters/cli/callback_test.go create mode 100644 apps/druid/adapters/cli/worker_progress_test.go create mode 100644 apps/druid/adapters/http/handlers/health_handler_test.go diff --git a/apps/druid/adapters/cli/callback.go b/apps/druid/adapters/cli/callback.go index cc154488..269cd4a0 100644 --- a/apps/druid/adapters/cli/callback.go +++ b/apps/druid/adapters/cli/callback.go @@ -11,6 +11,23 @@ type runtimeCallbackHandler struct { callbacks *appservices.WorkerCallbackManager } +func (h runtimeCallbackHandler) ReportProgress(c *fiber.Ctx) error { + var report struct { + Token string `json:"token"` + Percentage *int64 `json:"percentage"` + } + if err := c.BodyParser(&report); err != nil || report.Percentage == nil { + return fiber.NewError(fiber.StatusBadRequest, "invalid progress report") + } + if *report.Percentage < 0 || *report.Percentage > 100 { + return fiber.NewError(fiber.StatusBadRequest, "percentage must be between 0 and 100") + } + if err := h.callbacks.ReportProgress(c.Params("runtime_id"), report.Token, *report.Percentage); err != nil { + return fiber.NewError(fiber.StatusUnauthorized, err.Error()) + } + return c.SendStatus(fiber.StatusNoContent) +} + func (h runtimeCallbackHandler) CompleteWorker(c *fiber.Ctx, runtimeID callbackapi.Runtime) error { var result callbackapi.WorkerResult if err := c.BodyParser(&result); err != nil { diff --git a/apps/druid/adapters/cli/callback_test.go b/apps/druid/adapters/cli/callback_test.go new file mode 100644 index 00000000..259e11b6 --- /dev/null +++ b/apps/druid/adapters/cli/callback_test.go @@ -0,0 +1,40 @@ +package cli + +import ( + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gofiber/fiber/v2" + appservices "github.com/highcard-dev/daemon/apps/druid/core/services" +) + +func TestRuntimeCallbackHandlerReportsProgress(t *testing.T) { + callbacks := appservices.NewWorkerCallbackManager() + token, _, err := callbacks.Register("runtime-1") + if err != nil { + t.Fatal(err) + } + handler := runtimeCallbackHandler{callbacks: callbacks} + app := fiber.New() + app.Post("/internal/v1/workers/:runtime_id/progress", handler.ReportProgress) + + request := httptest.NewRequest( + http.MethodPost, + "/internal/v1/workers/runtime-1/progress", + strings.NewReader(fmt.Sprintf(`{"token":%q,"percentage":42}`, token)), + ) + request.Header.Set(fiber.HeaderContentType, fiber.MIMEApplicationJSON) + response, err := app.Test(request) + if err != nil { + t.Fatal(err) + } + if response.StatusCode != http.StatusNoContent { + t.Fatalf("status = %d; want %d", response.StatusCode, http.StatusNoContent) + } + if progress, ok := callbacks.Progress("runtime-1"); !ok || progress != 42 { + t.Fatalf("progress = %v, %v; want 42, true", progress, ok) + } +} diff --git a/apps/druid/adapters/cli/daemon.go b/apps/druid/adapters/cli/daemon.go index 7369d314..622d1426 100644 --- a/apps/druid/adapters/cli/daemon.go +++ b/apps/druid/adapters/cli/daemon.go @@ -170,7 +170,7 @@ func runRuntimeDaemon() error { websocketHandler.SetAllowUnauthenticatedPublic(runtimeAllowUnauthenticatedPublic) handlers := runtimehandlers.RouteHandlers{ Server: runtimehandlers.NewRuntimeServer( - runtimehandlers.NewHealthHandler(), + runtimehandlers.NewHealthHandlerWithProgress(callbacks.Progress), scrollHandler, ), Websocket: websocketHandler, @@ -206,7 +206,9 @@ func runRuntimeDaemon() error { if callbackListener != nil { callbackApp = fiber.New(fiber.Config{DisableStartupMessage: true, ErrorHandler: runtimehandlers.ErrorHandler}) callbackApp.Use(runtimehandlers.RequestLogger) - callbackapi.RegisterHandlers(callbackApp, runtimeCallbackHandler{callbacks: callbacks}) + callbackHandler := runtimeCallbackHandler{callbacks: callbacks} + callbackapi.RegisterHandlers(callbackApp, callbackHandler) + callbackApp.Post("/internal/v1/workers/:runtime_id/progress", callbackHandler.ReportProgress) } return listenRuntimeHTTP(managementApp, publicApp, callbackApp, callbackListener, runtime.Store.StateDir()) } diff --git a/apps/druid/adapters/cli/worker_progress_test.go b/apps/druid/adapters/cli/worker_progress_test.go new file mode 100644 index 00000000..42f7818c --- /dev/null +++ b/apps/druid/adapters/cli/worker_progress_test.go @@ -0,0 +1,53 @@ +package cli + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/highcard-dev/daemon/internal/core/domain" + "github.com/highcard-dev/daemon/internal/core/ports" +) + +func TestWorkerProgressReporterReadsSnapshotProgress(t *testing.T) { + reports := make(chan int64, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + var report struct { + Token string `json:"token"` + Percentage int64 `json:"percentage"` + } + if err := json.NewDecoder(request.Body).Decode(&report); err != nil { + t.Error(err) + } + if report.Token != "token" { + t.Errorf("token = %q; want token", report.Token) + } + reports <- report.Percentage + w.WriteHeader(http.StatusNoContent) + })) + defer server.Close() + + progress := domain.NewSnapshotProgress() + progress.Percentage.Store(37) + stop := startWorkerProgressReporter( + ports.RuntimeWorkerAction{ + RuntimeID: "runtime-1", + CallbackURL: server.URL + "/internal/v1/workers/runtime-1/complete", + CallbackToken: "token", + }, + progress, + time.Hour, + ) + defer stop() + + select { + case percentage := <-reports: + if percentage != 37 { + t.Fatalf("percentage = %d; want 37", percentage) + } + case <-time.After(time.Second): + t.Fatal("progress was not reported") + } +} diff --git a/apps/druid/adapters/cli/worker_pull.go b/apps/druid/adapters/cli/worker_pull.go index 1f7e9757..1a9ec26f 100644 --- a/apps/druid/adapters/cli/worker_pull.go +++ b/apps/druid/adapters/cli/worker_pull.go @@ -1,13 +1,16 @@ package cli import ( + "bytes" "context" "encoding/json" "fmt" "io" + "net/http" "os" "path/filepath" "strings" + "sync" "time" "github.com/highcard-dev/daemon/internal/callbackapi" @@ -65,6 +68,10 @@ func runWorkerPull(action ports.RuntimeWorkerAction) ports.RuntimeWorkerResult { if root == "" { root = "/scroll" } + progress := domain.NewSnapshotProgress() + stopProgress := startWorkerProgressReporter(action, progress, time.Second) + defer stopProgress() + oci := registry.NewOciClient(loadWorkerRegistryStore()) digest, err := oci.ResolveDigest(action.Artifact) if err == nil { @@ -72,11 +79,11 @@ func runWorkerPull(action ports.RuntimeWorkerAction) ports.RuntimeWorkerResult { } switch action.Mode { case ports.RuntimeWorkerModeUpdate: - err = pullWorkerUpdate(root, action.Artifact, oci) + err = pullWorkerUpdate(root, action.Artifact, oci, progress) case ports.RuntimeWorkerModeRestore: - err = pullWorkerRestore(root, action.Artifact, oci) + err = pullWorkerRestore(root, action.Artifact, oci, progress) default: - err = pullWorkerCreate(root, action.Artifact, oci) + err = pullWorkerCreate(root, action.Artifact, oci, progress) } if err != nil { result.Error = err.Error() @@ -115,7 +122,7 @@ func loadWorkerRegistryStore() *registry.CredentialStore { return registry.NewCredentialStore(config.Registries) } -func pullWorkerCreate(root string, artifact string, oci ports.OciRegistryInterface) error { +func pullWorkerCreate(root string, artifact string, oci ports.OciRegistryInterface, progress *domain.SnapshotProgress) error { if err := os.MkdirAll(root, 0755); err != nil { return err } @@ -137,16 +144,16 @@ func pullWorkerCreate(root string, artifact string, oci ports.OciRegistryInterfa } return copyPath(artifact, root) } - return oci.PullSelective(root, artifact, true, nil) + return oci.PullSelective(root, artifact, true, progress) } -func pullWorkerUpdate(root string, artifact string, oci ports.OciRegistryInterface) error { +func pullWorkerUpdate(root string, artifact string, oci ports.OciRegistryInterface, progress *domain.SnapshotProgress) error { tmp, err := os.MkdirTemp("", "druid-worker-update-*") if err != nil { return err } defer os.RemoveAll(tmp) - if err := coreservices.MaterializeScrollArtifact(artifact, tmp, oci, true); err != nil { + if err := coreservices.MaterializeScrollArtifactWithProgress(artifact, tmp, oci, true, progress); err != nil { return err } scrollYAML, err := os.ReadFile(filepath.Join(tmp, "scroll.yaml")) @@ -162,13 +169,13 @@ func pullWorkerUpdate(root string, artifact string, oci ports.OciRegistryInterfa return mergePulledRoot(tmp, root, skipData) } -func pullWorkerRestore(root string, artifact string, oci ports.OciRegistryInterface) error { +func pullWorkerRestore(root string, artifact string, oci ports.OciRegistryInterface, progress *domain.SnapshotProgress) error { tmp, err := os.MkdirTemp("", "druid-worker-restore-*") if err != nil { return err } defer os.RemoveAll(tmp) - if err := coreservices.MaterializeScrollArtifact(artifact, tmp, oci, true); err != nil { + if err := coreservices.MaterializeScrollArtifactWithProgress(artifact, tmp, oci, true, progress); err != nil { return err } if err := os.MkdirAll(root, 0755); err != nil { @@ -313,6 +320,62 @@ func copyPath(src string, dst string) error { return err } +func startWorkerProgressReporter(action ports.RuntimeWorkerAction, progress *domain.SnapshotProgress, interval time.Duration) func() { + if action.CallbackURL == "" || action.CallbackToken == "" || progress == nil { + return func() {} + } + suffix := "/internal/v1/workers/" + action.RuntimeID + "/complete" + baseURL := strings.TrimSuffix(action.CallbackURL, suffix) + if baseURL == action.CallbackURL { + return func() {} + } + progressURL := baseURL + "/internal/v1/workers/" + action.RuntimeID + "/progress" + client := &http.Client{Timeout: 3 * time.Second} + done := make(chan struct{}) + var wait sync.WaitGroup + wait.Add(1) + go func() { + defer wait.Done() + ticker := time.NewTicker(interval) + defer ticker.Stop() + lastPercentage := int64(-1) + report := func() { + percentage := progress.Percentage.Load() + if percentage == lastPercentage { + return + } + body, _ := json.Marshal(struct { + Token string `json:"token"` + Percentage int64 `json:"percentage"` + }{action.CallbackToken, percentage}) + request, _ := http.NewRequest(http.MethodPost, progressURL, bytes.NewReader(body)) + request.Header.Set("Content-Type", "application/json") + response, err := client.Do(request) + if err == nil { + response.Body.Close() + if response.StatusCode < http.StatusBadRequest { + lastPercentage = percentage + } + } + } + report() + for { + select { + case <-ticker.C: + report() + case <-done: + report() + return + } + } + }() + var once sync.Once + return func() { + once.Do(func() { close(done) }) + wait.Wait() + } +} + func reportWorkerResult(action ports.RuntimeWorkerAction, result ports.RuntimeWorkerResult) error { if action.CallbackURL == "" { body, err := json.Marshal(result) diff --git a/apps/druid/adapters/cli/worker_test.go b/apps/druid/adapters/cli/worker_test.go index 302145b4..5fe506c1 100644 --- a/apps/druid/adapters/cli/worker_test.go +++ b/apps/druid/adapters/cli/worker_test.go @@ -57,7 +57,7 @@ func TestWorkerRestoreStagesBeforeReplacingRoot(t *testing.T) { mustWrite(t, filepath.Join(root, "data", "old-only.txt"), "old") oci := fakeRestoreOCI{t: t} - if err := pullWorkerRestore(root, "registry.local/backup:1", oci); err != nil { + if err := pullWorkerRestore(root, "registry.local/backup:1", oci, nil); err != nil { t.Fatal(err) } diff --git a/apps/druid/adapters/http/handlers/health_handler.go b/apps/druid/adapters/http/handlers/health_handler.go index 0079f5e5..e7658d0a 100644 --- a/apps/druid/adapters/http/handlers/health_handler.go +++ b/apps/druid/adapters/http/handlers/health_handler.go @@ -5,12 +5,27 @@ import ( "github.com/highcard-dev/daemon/internal/api" ) -type HealthHandler struct{} +type ProgressLookup func(runtimeID string) (float64, bool) + +type HealthHandler struct { + progress ProgressLookup +} func NewHealthHandler() *HealthHandler { return &HealthHandler{} } +func NewHealthHandlerWithProgress(progress ProgressLookup) *HealthHandler { + return &HealthHandler{progress: progress} +} + func (h *HealthHandler) GetHealthAuth(c *fiber.Ctx) error { - return c.JSON(api.HealthResponse{Mode: "ok"}) + health := api.HealthResponse{Mode: "ok"} + if h.progress != nil { + if progress, ok := h.progress(c.Params("id")); ok { + value := float32(progress) + health.Progress = &value + } + } + return c.JSON(health) } diff --git a/apps/druid/adapters/http/handlers/health_handler_test.go b/apps/druid/adapters/http/handlers/health_handler_test.go new file mode 100644 index 00000000..8678e6b8 --- /dev/null +++ b/apps/druid/adapters/http/handlers/health_handler_test.go @@ -0,0 +1,33 @@ +package handlers + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gofiber/fiber/v2" + "github.com/highcard-dev/daemon/internal/api" +) + +func TestGetHealthAuthIncludesPullProgress(t *testing.T) { + handler := NewHealthHandlerWithProgress(func(runtimeID string) (float64, bool) { + return 37, runtimeID == "scroll-1" + }) + app := fiber.New() + app.Get("/:id/api/v1/health", handler.GetHealthAuth) + + response, err := app.Test(httptest.NewRequest(http.MethodGet, "/scroll-1/api/v1/health", nil)) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + + var health api.HealthResponse + if err := json.NewDecoder(response.Body).Decode(&health); err != nil { + t.Fatal(err) + } + if health.Progress == nil || *health.Progress != 37 { + t.Fatalf("progress = %v; want 37", health.Progress) + } +} diff --git a/apps/druid/core/services/worker_callbacks.go b/apps/druid/core/services/worker_callbacks.go index 69a27004..dba37751 100644 --- a/apps/druid/core/services/worker_callbacks.go +++ b/apps/druid/core/services/worker_callbacks.go @@ -10,8 +10,9 @@ import ( ) type WorkerCallbackManager struct { - mu sync.Mutex - actions map[string]workerCallbackAction + mu sync.Mutex + actions map[string]workerCallbackAction + progress map[string]int64 } type workerCallbackAction struct { @@ -20,7 +21,10 @@ type workerCallbackAction struct { } func NewWorkerCallbackManager() *WorkerCallbackManager { - return &WorkerCallbackManager{actions: map[string]workerCallbackAction{}} + return &WorkerCallbackManager{ + actions: map[string]workerCallbackAction{}, + progress: map[string]int64{}, + } } func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports.RuntimeWorkerResult, error) { @@ -36,6 +40,7 @@ func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports return "", nil, fmt.Errorf("worker action already pending for runtime %s", runtimeID) } m.actions[runtimeID] = workerCallbackAction{token: token, result: ch} + m.progress[runtimeID] = 0 m.mu.Unlock() return token, ch, nil } @@ -43,9 +48,31 @@ func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports func (m *WorkerCallbackManager) Cancel(runtimeID string) { m.mu.Lock() delete(m.actions, runtimeID) + delete(m.progress, runtimeID) m.mu.Unlock() } +func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, percentage int64) error { + m.mu.Lock() + defer m.mu.Unlock() + action, ok := m.actions[runtimeID] + if !ok { + return fmt.Errorf("unknown or completed worker action") + } + if token == "" || token != action.token { + return fmt.Errorf("invalid worker token") + } + m.progress[runtimeID] = max(0, min(100, percentage)) + return nil +} + +func (m *WorkerCallbackManager) Progress(runtimeID string) (float64, bool) { + m.mu.Lock() + defer m.mu.Unlock() + progress, ok := m.progress[runtimeID] + return float64(progress), ok +} + func (m *WorkerCallbackManager) Complete(runtimeID string, token string, result ports.RuntimeWorkerResult) error { m.mu.Lock() action, ok := m.actions[runtimeID] @@ -58,6 +85,7 @@ func (m *WorkerCallbackManager) Complete(runtimeID string, token string, result return fmt.Errorf("invalid worker token") } delete(m.actions, runtimeID) + delete(m.progress, runtimeID) m.mu.Unlock() action.result <- result close(action.result) diff --git a/apps/druid/core/services/worker_callbacks_test.go b/apps/druid/core/services/worker_callbacks_test.go index 2511b9ee..c9dd517e 100644 --- a/apps/druid/core/services/worker_callbacks_test.go +++ b/apps/druid/core/services/worker_callbacks_test.go @@ -53,3 +53,29 @@ func TestWorkerCallbackRejectsUnknownRuntime(t *testing.T) { t.Fatal("unknown runtime should fail") } } + +func TestWorkerCallbackTracksPullProgress(t *testing.T) { + manager := NewWorkerCallbackManager() + token, _, err := manager.Register("scroll-a") + if err != nil { + t.Fatal(err) + } + + if progress, ok := manager.Progress("scroll-a"); !ok || progress != 0 { + t.Fatalf("initial progress = %v, %v; want 0, true", progress, ok) + } + if err := manager.ReportProgress("scroll-a", "wrong-token", 42); err == nil { + t.Fatal("invalid progress token should fail") + } + if err := manager.ReportProgress("scroll-a", token, 42); err != nil { + t.Fatal(err) + } + if progress, ok := manager.Progress("scroll-a"); !ok || progress != 42 { + t.Fatalf("reported progress = %v, %v; want 42, true", progress, ok) + } + + manager.Cancel("scroll-a") + if _, ok := manager.Progress("scroll-a"); ok { + t.Fatal("cancelled progress should be removed") + } +} diff --git a/internal/core/services/registry/oci.go b/internal/core/services/registry/oci.go index 9d9dbdc6..226ebfda 100644 --- a/internal/core/services/registry/oci.go +++ b/internal/core/services/registry/oci.go @@ -201,6 +201,13 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, if progress != nil { progress.Mode.Store(domain.SnapshotProgressModeRestore) progress.Percentage.Store(0) + defer progress.Mode.Store(domain.SnapshotProgressModeIdle) + } + storeProgress := func(done, total int64) { + if progress == nil || total <= 0 { + return + } + progress.Percentage.Store(min(99, done*100/total)) } copyOpts := oras.CopyOptions{ @@ -259,10 +266,7 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, done := completed.Add(1) total := totalLayers.Load() bytesDownloaded.Add(desc.Size) - if progress != nil && total > 0 { - pct := done * 100 / total - progress.Percentage.Store(pct) - } + storeProgress(done, total) title := desc.Annotations["org.opencontainers.image.title"] logger.Log().Debug("Pulled layer", zap.String("title", title), @@ -279,10 +283,7 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, done := completed.Add(1) total := totalLayers.Load() bytesDownloaded.Add(desc.Size) - if progress != nil && total > 0 { - pct := done * 100 / total - progress.Percentage.Store(pct) - } + storeProgress(done, total) title := desc.Annotations["org.opencontainers.image.title"] logger.Log().Debug("Layer already exists locally, skipped", zap.String("title", title), @@ -303,17 +304,9 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, manifestDescriptor, err := oras.Copy(ctx, repoInstance, ref, fs, dstRef, copyOpts) stopProgress() if err != nil { - if progress != nil { - progress.Mode.Store(domain.SnapshotProgressModeIdle) - } return err } - if progress != nil { - progress.Percentage.Store(100) - progress.Mode.Store(domain.SnapshotProgressModeIdle) - } - logger.Log().Info("Manifest pulled", zap.String("digest", manifestDescriptor.Digest.String()), zap.String("mediaType", manifestDescriptor.MediaType)) jsonData, err := json.Marshal(&manifestDescriptor) @@ -349,6 +342,10 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, return fmt.Errorf("failed to write annotations: %w", err) } + if progress != nil { + progress.Percentage.Store(100) + } + return nil } diff --git a/internal/core/services/registry/oci_test.go b/internal/core/services/registry/oci_test.go index 2179b3d0..e8d5a841 100644 --- a/internal/core/services/registry/oci_test.go +++ b/internal/core/services/registry/oci_test.go @@ -117,6 +117,47 @@ func fakeRegistry(t *testing.T) *httptest.Server { return srv } +func TestPullSelectiveDoesNotReportCompleteBeforeMetadataIsWritten(t *testing.T) { + t.Chdir(t.TempDir()) + server := fakeRegistry(t) + registryHost := strings.TrimPrefix(server.URL, "http://") + source := filepath.Join("scrolls", "progress-test") + if err := os.MkdirAll(source, 0755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile( + filepath.Join(source, "scroll.yaml"), + []byte("name: progress-test\nversion: 0.1.0\napp_version: \"1.0\"\n"), + 0644, + ); err != nil { + t.Fatal(err) + } + + client := &OciClient{ + credentialStore: NewCredentialStore(nil), + plainHTTP: true, + } + repository := registryHost + "/test/progress" + if _, err := client.Push(source, repository, "1.0", nil, false, nil); err != nil { + t.Fatal(err) + } + + destination := filepath.Join("pull", "progress-test") + if err := os.MkdirAll(filepath.Join(destination, "manifest.json"), 0755); err != nil { + t.Fatal(err) + } + progress := domain.NewSnapshotProgress() + if err := client.PullSelective(destination, repository+":1.0", true, progress); err == nil { + t.Fatal("pull should fail when manifest.json cannot be written") + } + if percentage := progress.Percentage.Load(); percentage >= 100 { + t.Fatalf("failed pull progress = %d; want below 100", percentage) + } + if mode := progress.Mode.Load(); mode != domain.SnapshotProgressModeIdle { + t.Fatalf("failed pull mode = %v; want idle", mode) + } +} + func TestValidateCredentialsUsesPlainHTTPEnv(t *testing.T) { t.Setenv("DRUID_REGISTRY_PLAIN_HTTP", "true") diff --git a/internal/core/services/runtime_scroll_manager.go b/internal/core/services/runtime_scroll_manager.go index a431cb1a..ce0aba53 100644 --- a/internal/core/services/runtime_scroll_manager.go +++ b/internal/core/services/runtime_scroll_manager.go @@ -102,6 +102,10 @@ func RuntimeScrollIDFromName(name string) string { } func MaterializeScrollArtifact(artifact string, root string, ociRegistry ports.OciRegistryInterface, includeData bool) error { + return MaterializeScrollArtifactWithProgress(artifact, root, ociRegistry, includeData, nil) +} + +func MaterializeScrollArtifactWithProgress(artifact string, root string, ociRegistry ports.OciRegistryInterface, includeData bool, progress *domain.SnapshotProgress) error { if artifact == "" { return fmt.Errorf("artifact is required") } @@ -126,7 +130,7 @@ func MaterializeScrollArtifact(artifact string, root string, ociRegistry ports.O if ociRegistry == nil { return fmt.Errorf("OCI registry is required to pull %s", artifact) } - if err := ociRegistry.PullSelective(root, artifact, includeData, nil); err != nil { + if err := ociRegistry.PullSelective(root, artifact, includeData, progress); err != nil { return err } return os.MkdirAll(filepath.Join(root, domain.RuntimeDataDir), 0755) From d8a8e13c881ff9704bff7b694de8ae676528a4f2 Mon Sep 17 00:00:00 2001 From: IOHelpMe <101821466+IOHelpMe@users.noreply.github.com> Date: Fri, 24 Jul 2026 14:22:08 +0200 Subject: [PATCH 2/4] fix(runtime): keep SnapshotProgress current during setup --- apps/druid/core/services/runtime_lifecycle.go | 6 ++ apps/druid/core/services/runtime_session.go | 1 + .../core/services/runtime_session_cache.go | 6 ++ .../services/runtime_session_execution.go | 15 ++-- apps/druid/core/services/runtime_update.go | 3 + apps/druid/core/services/worker_callbacks.go | 60 ++++++++++++-- .../core/services/worker_callbacks_test.go | 55 +++++++++++++ internal/core/ports/services_ports.go | 1 + internal/runtime/kubernetes/procedures.go | 8 +- .../runtime/kubernetes/snapshot_progress.go | 61 ++++++++++++++ .../kubernetes/snapshot_progress_test.go | 81 +++++++++++++++++++ internal/runtime/kubernetes/wait_jobs.go | 75 ++++++++++++----- 12 files changed, 336 insertions(+), 36 deletions(-) create mode 100644 internal/runtime/kubernetes/snapshot_progress.go create mode 100644 internal/runtime/kubernetes/snapshot_progress_test.go diff --git a/apps/druid/core/services/runtime_lifecycle.go b/apps/druid/core/services/runtime_lifecycle.go index a0df0ef3..9c7730d4 100644 --- a/apps/druid/core/services/runtime_lifecycle.go +++ b/apps/druid/core/services/runtime_lifecycle.go @@ -13,6 +13,9 @@ func (s *RuntimeSupervisor) DeleteWithPolicy(id string, purgeData bool) error { s.mu.Unlock() if session != nil { session.stopDeploymentQueue() + if s.workerCallbacks != nil { + s.workerCallbacks.ClearSnapshotProgress(id, session.snapshotProgress) + } } runtimeScroll, err := s.store.GetScroll(id) @@ -60,6 +63,9 @@ func (s *RuntimeSupervisor) Stop(id string) (*domain.RuntimeScroll, error) { session.markError(err) return nil, err } + if s.workerCallbacks != nil { + s.workerCallbacks.ClearSnapshotProgress(id, session.snapshotProgress) + } session.stopDeploymentQueue() return s.store.GetScroll(id) } diff --git a/apps/druid/core/services/runtime_session.go b/apps/druid/core/services/runtime_session.go index 83e31d25..4c140129 100644 --- a/apps/druid/core/services/runtime_session.go +++ b/apps/druid/core/services/runtime_session.go @@ -19,6 +19,7 @@ type RuntimeSession struct { scrollService *coreservices.ScrollService watchService ports.WatchServiceInterface runtimeBackend ports.RuntimeBackendInterface + snapshotProgress *domain.SnapshotProgress queue map[string]*runtimeQueueItem workWg sync.WaitGroup notifierChan []chan []string diff --git a/apps/druid/core/services/runtime_session_cache.go b/apps/druid/core/services/runtime_session_cache.go index 8b4063e2..669586e4 100644 --- a/apps/druid/core/services/runtime_session_cache.go +++ b/apps/druid/core/services/runtime_session_cache.go @@ -54,12 +54,18 @@ func (s *RuntimeSupervisor) startSession(runtimeScroll *domain.RuntimeScroll) (* session.devDaemonToken = s.internalToken session.devAuthJWKSURL = s.authJWKSURL session.devRuntimeJWKSURL = s.runtimeJWKSURL + if s.workerCallbacks != nil { + session.snapshotProgress = s.workerCallbacks.TrackSnapshotProgress(runtimeScroll.ID) + } session.Start() s.mu.Lock() if existing := s.sessions[runtimeScroll.ID]; existing != nil { s.mu.Unlock() session.stopDeploymentQueue() + if s.workerCallbacks != nil { + s.workerCallbacks.ClearSnapshotProgress(runtimeScroll.ID, session.snapshotProgress) + } return existing, nil } s.sessions[runtimeScroll.ID] = session diff --git a/apps/druid/core/services/runtime_session_execution.go b/apps/druid/core/services/runtime_session_execution.go index b845c91c..5c7961a5 100644 --- a/apps/druid/core/services/runtime_session_execution.go +++ b/apps/druid/core/services/runtime_session_execution.go @@ -52,13 +52,14 @@ func (s *RuntimeSession) runCommand(cmd string) error { } exitCode, err := s.runtimeBackend.RunCommand(ports.RuntimeCommand{ - Name: cmd, - ScrollID: scrollID, - Command: command, - Root: root, - GlobalPorts: runtimePorts, - Routing: routing, - ProcedureEnv: procedureEnv, + Name: cmd, + ScrollID: scrollID, + Command: command, + Root: root, + GlobalPorts: runtimePorts, + Routing: routing, + ProcedureEnv: procedureEnv, + SnapshotProgress: s.snapshotProgress, ProcedureStatusObserver: func(procedure string, status domain.ScrollLockStatus, exitCode *int) { s.persistProcedureStatus(cmd, procedure, status, exitCode) }, diff --git a/apps/druid/core/services/runtime_update.go b/apps/druid/core/services/runtime_update.go index 03e22bf9..92190953 100644 --- a/apps/druid/core/services/runtime_update.go +++ b/apps/druid/core/services/runtime_update.go @@ -33,6 +33,9 @@ func (s *RuntimeSupervisor) updateExistingScroll(runtimeScroll *domain.RuntimeSc s.mu.Unlock() if session != nil { session.stopDeploymentQueue() + if s.workerCallbacks != nil { + s.workerCallbacks.ClearSnapshotProgress(runtimeScroll.ID, session.snapshotProgress) + } } if wasRunning { diff --git a/apps/druid/core/services/worker_callbacks.go b/apps/druid/core/services/worker_callbacks.go index dba37751..f6ecc475 100644 --- a/apps/druid/core/services/worker_callbacks.go +++ b/apps/druid/core/services/worker_callbacks.go @@ -6,13 +6,14 @@ import ( "fmt" "sync" + "github.com/highcard-dev/daemon/internal/core/domain" "github.com/highcard-dev/daemon/internal/core/ports" ) type WorkerCallbackManager struct { mu sync.Mutex actions map[string]workerCallbackAction - progress map[string]int64 + progress map[string]workerCallbackProgress } type workerCallbackAction struct { @@ -20,10 +21,15 @@ type workerCallbackAction struct { result chan ports.RuntimeWorkerResult } +type workerCallbackProgress struct { + snapshot *domain.SnapshotProgress + trackers int +} + func NewWorkerCallbackManager() *WorkerCallbackManager { return &WorkerCallbackManager{ actions: map[string]workerCallbackAction{}, - progress: map[string]int64{}, + progress: map[string]workerCallbackProgress{}, } } @@ -40,7 +46,12 @@ func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports return "", nil, fmt.Errorf("worker action already pending for runtime %s", runtimeID) } m.actions[runtimeID] = workerCallbackAction{token: token, result: ch} - m.progress[runtimeID] = 0 + progress := m.progress[runtimeID] + if progress.snapshot == nil { + progress.snapshot = domain.NewSnapshotProgress() + } + progress.snapshot.Percentage.Store(0) + m.progress[runtimeID] = progress m.mu.Unlock() return token, ch, nil } @@ -48,7 +59,9 @@ func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports func (m *WorkerCallbackManager) Cancel(runtimeID string) { m.mu.Lock() delete(m.actions, runtimeID) - delete(m.progress, runtimeID) + if progress := m.progress[runtimeID]; progress.trackers == 0 { + delete(m.progress, runtimeID) + } m.mu.Unlock() } @@ -62,7 +75,7 @@ func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, p if token == "" || token != action.token { return fmt.Errorf("invalid worker token") } - m.progress[runtimeID] = max(0, min(100, percentage)) + m.progress[runtimeID].snapshot.Percentage.Store(max(0, min(100, percentage))) return nil } @@ -70,7 +83,38 @@ func (m *WorkerCallbackManager) Progress(runtimeID string) (float64, bool) { m.mu.Lock() defer m.mu.Unlock() progress, ok := m.progress[runtimeID] - return float64(progress), ok + if !ok { + return 0, false + } + return float64(progress.snapshot.Percentage.Load()), true +} + +func (m *WorkerCallbackManager) TrackSnapshotProgress(runtimeID string) *domain.SnapshotProgress { + m.mu.Lock() + progress := m.progress[runtimeID] + if progress.snapshot == nil { + progress.snapshot = domain.NewSnapshotProgress() + } + progress.trackers++ + m.progress[runtimeID] = progress + m.mu.Unlock() + return progress.snapshot +} + +func (m *WorkerCallbackManager) ClearSnapshotProgress(runtimeID string, progress *domain.SnapshotProgress) { + m.mu.Lock() + current, ok := m.progress[runtimeID] + if ok && current.snapshot == progress { + if current.trackers > 0 { + current.trackers-- + } + if _, pending := m.actions[runtimeID]; pending || current.trackers > 0 { + m.progress[runtimeID] = current + } else { + delete(m.progress, runtimeID) + } + } + m.mu.Unlock() } func (m *WorkerCallbackManager) Complete(runtimeID string, token string, result ports.RuntimeWorkerResult) error { @@ -85,7 +129,9 @@ func (m *WorkerCallbackManager) Complete(runtimeID string, token string, result return fmt.Errorf("invalid worker token") } delete(m.actions, runtimeID) - delete(m.progress, runtimeID) + if progress := m.progress[runtimeID]; progress.trackers == 0 { + delete(m.progress, runtimeID) + } m.mu.Unlock() action.result <- result close(action.result) diff --git a/apps/druid/core/services/worker_callbacks_test.go b/apps/druid/core/services/worker_callbacks_test.go index c9dd517e..37157f99 100644 --- a/apps/druid/core/services/worker_callbacks_test.go +++ b/apps/druid/core/services/worker_callbacks_test.go @@ -79,3 +79,58 @@ func TestWorkerCallbackTracksPullProgress(t *testing.T) { t.Fatal("cancelled progress should be removed") } } + +func TestWorkerCallbackReadsTrackedSnapshotProgress(t *testing.T) { + manager := NewWorkerCallbackManager() + progress := manager.TrackSnapshotProgress("scroll-a") + if again := manager.TrackSnapshotProgress("scroll-a"); again != progress { + t.Fatal("tracking the same runtime replaced SnapshotProgress") + } + progress.Percentage.Store(43) + + if got, ok := manager.Progress("scroll-a"); !ok || got != 43 { + t.Fatalf("progress = %v, %v; want 43, true", got, ok) + } + + manager.ClearSnapshotProgress("scroll-a", progress) + if _, ok := manager.Progress("scroll-a"); !ok { + t.Fatal("clearing one tracker removed another tracker's SnapshotProgress") + } + manager.ClearSnapshotProgress("scroll-a", progress) + if _, ok := manager.Progress("scroll-a"); ok { + t.Fatal("cleared SnapshotProgress should be removed") + } +} + +func TestWorkerCallbackKeepsTrackedSnapshotAcrossWorkerLifecycle(t *testing.T) { + manager := NewWorkerCallbackManager() + progress := manager.TrackSnapshotProgress("scroll-a") + token, _, err := manager.Register("scroll-a") + if err != nil { + t.Fatal(err) + } + manager.mu.Lock() + registered := manager.progress["scroll-a"].snapshot + manager.mu.Unlock() + if registered != progress { + t.Fatal("worker registration replaced tracked SnapshotProgress") + } + if err := manager.ReportProgress("scroll-a", token, 42); err != nil { + t.Fatal(err) + } + if got := progress.Percentage.Load(); got != 42 { + t.Fatalf("tracked percentage = %d; want 42", got) + } + if err := manager.Complete("scroll-a", token, ports.RuntimeWorkerResult{}); err != nil { + t.Fatal(err) + } + + progress.Percentage.Store(43) + if got, ok := manager.Progress("scroll-a"); !ok || got != 43 { + t.Fatalf("progress after worker completion = %v, %v; want 43, true", got, ok) + } + manager.ClearSnapshotProgress("scroll-a", progress) + if _, ok := manager.Progress("scroll-a"); ok { + t.Fatal("cleared tracked progress should be removed") + } +} diff --git a/internal/core/ports/services_ports.go b/internal/core/ports/services_ports.go index 5de1eb3f..cd0a642f 100644 --- a/internal/core/ports/services_ports.go +++ b/internal/core/ports/services_ports.go @@ -80,6 +80,7 @@ type RuntimeCommand struct { GlobalPorts []domain.Port Routing []domain.RuntimeRouteAssignment ProcedureEnv map[string]map[string]string + SnapshotProgress *domain.SnapshotProgress ProcedureStatusObserver func(procedure string, status domain.ScrollLockStatus, exitCode *int) } diff --git a/internal/runtime/kubernetes/procedures.go b/internal/runtime/kubernetes/procedures.go index 6b02c2d9..0cb745d7 100644 --- a/internal/runtime/kubernetes/procedures.go +++ b/internal/runtime/kubernetes/procedures.go @@ -108,7 +108,7 @@ func (b *Backend) RunCommand(command ports.RuntimeCommand) (*int, error) { continue } command.ObserveProcedureStatus(procedureName, domain.ScrollLockStatusRunning, nil) - exitCode, err := b.runJobProcedure(command.ScrollID, command.Name, procedureName, resourceName, procedure, command.Root, command.GlobalPorts, env, portUse) + exitCode, err := b.runJobProcedure(command.ScrollID, command.Name, procedureName, resourceName, procedure, command.Root, command.GlobalPorts, env, portUse, command.SnapshotProgress) if err != nil { if exitCode != nil && *exitCode != 0 && procedure.IgnoreFailure { command.ObserveProcedureStatus(procedureName, domain.ScrollLockStatusDone, exitCode) @@ -138,7 +138,7 @@ func (b *Backend) RunCommand(command ports.RuntimeCommand) (*int, error) { return nil, nil } -func (b *Backend) runJobProcedure(scrollID string, commandName string, procedureName string, resourceName string, procedure *domain.Procedure, root string, globalPorts []domain.Port, env map[string]string, portUse map[string]int) (*int, error) { +func (b *Backend) runJobProcedure(scrollID string, commandName string, procedureName string, resourceName string, procedure *domain.Procedure, root string, globalPorts []domain.Port, env map[string]string, portUse map[string]int, progress *domain.SnapshotProgress) (*int, error) { if procedure.IsSignal() { logger.Log().Info("Running Kubernetes signal procedure", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("target", procedure.Target), zap.String("signal", procedure.Signal)) if err := b.Signal(procedureName, procedure.Target, procedure.Signal, root); err != nil { @@ -196,7 +196,7 @@ func (b *Backend) runJobProcedure(scrollID string, commandName string, procedure if err == nil { streamStarted = true logger.Log().Debug("Streaming Kubernetes job procedure logs", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("namespace", namespace), zap.String("job", jobName), zap.String("pod", podName), zap.String("console_id", consoleID)) - go b.streamPodLogs(ctx, namespace, podName, output) + go b.streamPodLogs(ctx, namespace, podName, output, progress) } else { logger.Log().Warn("Could not find Kubernetes job pod before wait; console logs may be empty", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("namespace", namespace), zap.String("job", jobName), zap.Error(err)) } @@ -301,7 +301,7 @@ func (b *Backend) ensurePersistentProcedure(ctx context.Context, scrollID string return } logger.Log().Debug("Streaming Kubernetes persistent procedure logs", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("namespace", namespace), zap.String("pod", podName)) - b.streamPodLogs(context.Background(), namespace, podName, output) + b.streamPodLogs(context.Background(), namespace, podName, output, nil) }() return nil } diff --git a/internal/runtime/kubernetes/snapshot_progress.go b/internal/runtime/kubernetes/snapshot_progress.go new file mode 100644 index 00000000..7b70a985 --- /dev/null +++ b/internal/runtime/kubernetes/snapshot_progress.go @@ -0,0 +1,61 @@ +package kubernetes + +import ( + "encoding/json" + "math" + "regexp" + "strconv" + "strings" + + "github.com/highcard-dev/daemon/internal/core/domain" +) + +const snapshotProgressPrefix = "DRUID_PROGRESS_V1 " + +var steamCMDProgressPattern = regexp.MustCompile( + `Update state \(0x[0-9a-fA-F]+\) downloading, progress: [0-9]+(?:\.[0-9]+)? \(([0-9]+) / ([0-9]+)\)`, +) + +type snapshotProgressSample struct { + Unit string `json:"unit"` + Current float64 `json:"current"` + Total float64 `json:"total"` +} + +func observeSnapshotProgress(line string, progress *domain.SnapshotProgress) bool { + if progress == nil { + return false + } + + if strings.HasPrefix(line, snapshotProgressPrefix) { + var sample snapshotProgressSample + if err := json.Unmarshal([]byte(strings.TrimPrefix(line, snapshotProgressPrefix)), &sample); err != nil { + return false + } + if sample.Unit != "bytes" || !storeSnapshotProgress(progress, sample.Current, sample.Total) { + return false + } + return true + } + + matches := steamCMDProgressPattern.FindStringSubmatch(line) + if len(matches) == 3 { + current, currentErr := strconv.ParseFloat(matches[1], 64) + total, totalErr := strconv.ParseFloat(matches[2], 64) + if currentErr == nil && totalErr == nil { + storeSnapshotProgress(progress, current, total) + } + } + return false +} + +func storeSnapshotProgress(progress *domain.SnapshotProgress, current float64, total float64) bool { + if current < 0 || total <= 0 || + math.IsNaN(current) || math.IsInf(current, 0) || + math.IsNaN(total) || math.IsInf(total, 0) { + return false + } + percentage := int64(math.Round(current / total * 100)) + progress.Percentage.Store(max(0, min(100, percentage))) + return true +} diff --git a/internal/runtime/kubernetes/snapshot_progress_test.go b/internal/runtime/kubernetes/snapshot_progress_test.go new file mode 100644 index 00000000..ff4125d3 --- /dev/null +++ b/internal/runtime/kubernetes/snapshot_progress_test.go @@ -0,0 +1,81 @@ +package kubernetes + +import ( + "strings" + "testing" + "time" + + "github.com/highcard-dev/daemon/internal/core/domain" +) + +func TestObserveSnapshotProgressReadsSteamCMDByteMarker(t *testing.T) { + progress := domain.NewSnapshotProgress() + line := `DRUID_PROGRESS_V1 {"step_id":"steamcmd","stage":"downloading","label":"Downloading server files","unit":"bytes","current":12549178426,"total":22938933947,"total_final":true}` + + if !observeSnapshotProgress(line, progress) { + t.Fatal("progress marker was not recognized") + } + if got := progress.Percentage.Load(); got != 55 { + t.Fatalf("percentage = %d; want 55", got) + } +} + +func TestObserveSnapshotProgressReadsOriginalSteamCMDLine(t *testing.T) { + progress := domain.NewSnapshotProgress() + line := "\x1b[0m Update state (0x61) downloading, progress: 54.71 (12549178426 / 22938933947)" + + if observeSnapshotProgress(line, progress) { + t.Fatal("original SteamCMD output should remain visible in the console") + } + if got := progress.Percentage.Load(); got != 55 { + t.Fatalf("percentage = %d; want 55", got) + } +} + +func TestReadSnapshotProgressDoesNotDependOnConsoleConsumer(t *testing.T) { + output := make(chan string, 1) + output <- "older console line" + progress := domain.NewSnapshotProgress() + done := make(chan struct{}) + + go func() { + _ = readSnapshotProgress( + strings.NewReader( + "ordinary console line\n"+ + `DRUID_PROGRESS_V1 {"unit":"bytes","current":12549178426,"total":22938933947}`+"\n", + ), + progress, + ) + close(done) + }() + + select { + case <-done: + case <-time.After(100 * time.Millisecond): + t.Fatal("progress processing blocked behind the console backlog") + } + if got := progress.Percentage.Load(); got != 55 { + t.Fatalf("percentage = %d; want 55", got) + } + if got := <-output; got != "older console line" { + t.Fatalf("console backlog changed to %q", got) + } +} + +func TestObserveSnapshotProgressLeavesMalformedLinesAlone(t *testing.T) { + progress := domain.NewSnapshotProgress() + progress.Percentage.Store(17) + + for _, line := range []string{ + "ordinary server output", + `DRUID_PROGRESS_V1 {"unit":"bytes","current":10,"total":0}`, + `DRUID_PROGRESS_V1 {"unit":"items","current":10,"total":20}`, + } { + if observeSnapshotProgress(line, progress) { + t.Fatalf("line was unexpectedly recognized: %s", line) + } + if got := progress.Percentage.Load(); got != 17 { + t.Fatalf("percentage changed to %d for line %q", got, line) + } + } +} diff --git a/internal/runtime/kubernetes/wait_jobs.go b/internal/runtime/kubernetes/wait_jobs.go index 5e1422d1..71436e53 100644 --- a/internal/runtime/kubernetes/wait_jobs.go +++ b/internal/runtime/kubernetes/wait_jobs.go @@ -14,6 +14,7 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/labels" + "github.com/highcard-dev/daemon/internal/core/domain" "github.com/highcard-dev/daemon/internal/utils/logger" "go.uber.org/zap" ) @@ -253,48 +254,86 @@ func (b *Backend) podLogs(ctx context.Context, namespace string, podName string) return logs, nil } -func (b *Backend) streamPodLogs(ctx context.Context, namespace string, podName string, output chan<- string) { +func (b *Backend) streamPodLogs(ctx context.Context, namespace string, podName string, output chan<- string, progress *domain.SnapshotProgress) { defer close(output) - var stream io.ReadCloser + stream, err := b.openPodLogStream(ctx, namespace, podName) + if err != nil { + output <- fmt.Sprintf("failed to stream pod logs: %v", err) + return + } + defer stream.Close() + var progressDone chan struct{} + if progress != nil { + progressDone = make(chan struct{}) + go func() { + defer close(progressDone) + b.streamPodProgress(ctx, namespace, podName, progress) + }() + } + scanner := bufio.NewScanner(stream) + for scanner.Scan() { + line := scanner.Text() + if progress != nil && strings.HasPrefix(line, snapshotProgressPrefix) { + continue + } + output <- line + } + if progressDone != nil { + <-progressDone + } + if err := scanner.Err(); err != nil { + logger.Log().Warn("Kubernetes pod log stream ended with scanner error", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) + return + } + logger.Log().Debug("Kubernetes pod log stream ended", zap.String("namespace", namespace), zap.String("pod", podName)) +} + +func (b *Backend) openPodLogStream(ctx context.Context, namespace string, podName string) (io.ReadCloser, error) { deadline := time.Now().Add(30 * time.Second) logger.Log().Debug("Opening Kubernetes follow log stream", zap.String("namespace", namespace), zap.String("pod", podName)) for { req := b.client.CoreV1().Pods(namespace).GetLogs(podName, &corev1.PodLogOptions{Follow: true}) - var err error - stream, err = req.Stream(ctx) + stream, err := req.Stream(ctx) if err == nil { logger.Log().Debug("Kubernetes follow log stream opened", zap.String("namespace", namespace), zap.String("pod", podName)) - break + return stream, nil } if !strings.Contains(err.Error(), "ContainerCreating") && !strings.Contains(err.Error(), "PodInitializing") && !strings.Contains(err.Error(), "not available") { logger.Log().Warn("Failed to stream Kubernetes pod logs", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - output <- fmt.Sprintf("failed to stream pod logs: %v", err) - return + return nil, err } if time.Now().After(deadline) { logger.Log().Warn("Timed out opening Kubernetes pod log stream", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - output <- fmt.Sprintf("failed to stream pod logs: %v", err) - return + return nil, err } logger.Log().Debug("Kubernetes pod logs not ready yet", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) select { case <-ctx.Done(): logger.Log().Warn("Context cancelled while opening Kubernetes pod logs", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(ctx.Err())) - output <- fmt.Sprintf("failed to stream pod logs: %v", ctx.Err()) - return + return nil, ctx.Err() case <-time.After(500 * time.Millisecond): } } +} + +func (b *Backend) streamPodProgress(ctx context.Context, namespace string, podName string, progress *domain.SnapshotProgress) { + stream, err := b.openPodLogStream(ctx, namespace, podName) + if err != nil { + logger.Log().Warn("Failed to stream Kubernetes pod progress", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) + return + } defer stream.Close() - scanner := bufio.NewScanner(stream) - for scanner.Scan() { - output <- scanner.Text() + if err := readSnapshotProgress(stream, progress); err != nil { + logger.Log().Warn("Kubernetes pod progress stream ended with scanner error", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) } - if err := scanner.Err(); err != nil { - logger.Log().Warn("Kubernetes pod log stream ended with scanner error", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - return +} + +func readSnapshotProgress(input io.Reader, progress *domain.SnapshotProgress) error { + scanner := bufio.NewScanner(input) + for scanner.Scan() { + observeSnapshotProgress(scanner.Text(), progress) } - logger.Log().Debug("Kubernetes pod log stream ended", zap.String("namespace", namespace), zap.String("pod", podName)) + return scanner.Err() } From 30903887044f8b2692c74531a18bab30cba18440 Mon Sep 17 00:00:00 2001 From: IOHelpMe <101821466+IOHelpMe@users.noreply.github.com> Date: Sat, 25 Jul 2026 15:33:10 +0200 Subject: [PATCH 3/4] fix(runtime): preserve SnapshotProgress tenths --- apps/druid/adapters/cli/callback.go | 11 ++++--- apps/druid/adapters/cli/callback_test.go | 6 ++-- .../adapters/cli/worker_progress_test.go | 12 +++---- apps/druid/adapters/cli/worker_pull.go | 8 ++--- apps/druid/core/services/worker_callbacks.go | 10 +++--- .../core/services/worker_callbacks_test.go | 26 ++++++++-------- internal/core/domain/oci.go | 22 +++++++++++-- internal/core/domain/oci_test.go | 31 +++++++++++++++++++ .../coldstarter/handler/lua_handler.go | 2 +- internal/core/services/registry/oci.go | 7 +++-- internal/core/services/registry/oci_test.go | 4 +-- .../runtime/kubernetes/snapshot_progress.go | 4 +-- .../kubernetes/snapshot_progress_test.go | 30 ++++++++++++------ 13 files changed, 119 insertions(+), 54 deletions(-) create mode 100644 internal/core/domain/oci_test.go diff --git a/apps/druid/adapters/cli/callback.go b/apps/druid/adapters/cli/callback.go index 269cd4a0..8a9ab3af 100644 --- a/apps/druid/adapters/cli/callback.go +++ b/apps/druid/adapters/cli/callback.go @@ -1,6 +1,8 @@ package cli import ( + "math" + "github.com/gofiber/fiber/v2" appservices "github.com/highcard-dev/daemon/apps/druid/core/services" "github.com/highcard-dev/daemon/internal/callbackapi" @@ -13,16 +15,17 @@ type runtimeCallbackHandler struct { func (h runtimeCallbackHandler) ReportProgress(c *fiber.Ctx) error { var report struct { - Token string `json:"token"` - Percentage *int64 `json:"percentage"` + Token string `json:"token"` + Percentage *float64 `json:"percentage"` } if err := c.BodyParser(&report); err != nil || report.Percentage == nil { return fiber.NewError(fiber.StatusBadRequest, "invalid progress report") } - if *report.Percentage < 0 || *report.Percentage > 100 { + percentage := *report.Percentage + if math.IsNaN(percentage) || math.IsInf(percentage, 0) || percentage < 0 || percentage > 100 { return fiber.NewError(fiber.StatusBadRequest, "percentage must be between 0 and 100") } - if err := h.callbacks.ReportProgress(c.Params("runtime_id"), report.Token, *report.Percentage); err != nil { + if err := h.callbacks.ReportProgress(c.Params("runtime_id"), report.Token, percentage); err != nil { return fiber.NewError(fiber.StatusUnauthorized, err.Error()) } return c.SendStatus(fiber.StatusNoContent) diff --git a/apps/druid/adapters/cli/callback_test.go b/apps/druid/adapters/cli/callback_test.go index 259e11b6..0b9465bb 100644 --- a/apps/druid/adapters/cli/callback_test.go +++ b/apps/druid/adapters/cli/callback_test.go @@ -24,7 +24,7 @@ func TestRuntimeCallbackHandlerReportsProgress(t *testing.T) { request := httptest.NewRequest( http.MethodPost, "/internal/v1/workers/runtime-1/progress", - strings.NewReader(fmt.Sprintf(`{"token":%q,"percentage":42}`, token)), + strings.NewReader(fmt.Sprintf(`{"token":%q,"percentage":42.3}`, token)), ) request.Header.Set(fiber.HeaderContentType, fiber.MIMEApplicationJSON) response, err := app.Test(request) @@ -34,7 +34,7 @@ func TestRuntimeCallbackHandlerReportsProgress(t *testing.T) { if response.StatusCode != http.StatusNoContent { t.Fatalf("status = %d; want %d", response.StatusCode, http.StatusNoContent) } - if progress, ok := callbacks.Progress("runtime-1"); !ok || progress != 42 { - t.Fatalf("progress = %v, %v; want 42, true", progress, ok) + if progress, ok := callbacks.Progress("runtime-1"); !ok || progress != 42.3 { + t.Fatalf("progress = %v, %v; want 42.3, true", progress, ok) } } diff --git a/apps/druid/adapters/cli/worker_progress_test.go b/apps/druid/adapters/cli/worker_progress_test.go index 42f7818c..c5891e92 100644 --- a/apps/druid/adapters/cli/worker_progress_test.go +++ b/apps/druid/adapters/cli/worker_progress_test.go @@ -12,11 +12,11 @@ import ( ) func TestWorkerProgressReporterReadsSnapshotProgress(t *testing.T) { - reports := make(chan int64, 1) + reports := make(chan float64, 1) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { var report struct { - Token string `json:"token"` - Percentage int64 `json:"percentage"` + Token string `json:"token"` + Percentage float64 `json:"percentage"` } if err := json.NewDecoder(request.Body).Decode(&report); err != nil { t.Error(err) @@ -30,7 +30,7 @@ func TestWorkerProgressReporterReadsSnapshotProgress(t *testing.T) { defer server.Close() progress := domain.NewSnapshotProgress() - progress.Percentage.Store(37) + progress.StorePercentage(37.4) stop := startWorkerProgressReporter( ports.RuntimeWorkerAction{ RuntimeID: "runtime-1", @@ -44,8 +44,8 @@ func TestWorkerProgressReporterReadsSnapshotProgress(t *testing.T) { select { case percentage := <-reports: - if percentage != 37 { - t.Fatalf("percentage = %d; want 37", percentage) + if percentage != 37.4 { + t.Fatalf("percentage = %v; want 37.4", percentage) } case <-time.After(time.Second): t.Fatal("progress was not reported") diff --git a/apps/druid/adapters/cli/worker_pull.go b/apps/druid/adapters/cli/worker_pull.go index 1a9ec26f..dd0da557 100644 --- a/apps/druid/adapters/cli/worker_pull.go +++ b/apps/druid/adapters/cli/worker_pull.go @@ -338,15 +338,15 @@ func startWorkerProgressReporter(action ports.RuntimeWorkerAction, progress *dom defer wait.Done() ticker := time.NewTicker(interval) defer ticker.Stop() - lastPercentage := int64(-1) + lastPercentage := -1.0 report := func() { - percentage := progress.Percentage.Load() + percentage := progress.Percentage() if percentage == lastPercentage { return } body, _ := json.Marshal(struct { - Token string `json:"token"` - Percentage int64 `json:"percentage"` + Token string `json:"token"` + Percentage float64 `json:"percentage"` }{action.CallbackToken, percentage}) request, _ := http.NewRequest(http.MethodPost, progressURL, bytes.NewReader(body)) request.Header.Set("Content-Type", "application/json") diff --git a/apps/druid/core/services/worker_callbacks.go b/apps/druid/core/services/worker_callbacks.go index f6ecc475..470c6152 100644 --- a/apps/druid/core/services/worker_callbacks.go +++ b/apps/druid/core/services/worker_callbacks.go @@ -50,7 +50,7 @@ func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports if progress.snapshot == nil { progress.snapshot = domain.NewSnapshotProgress() } - progress.snapshot.Percentage.Store(0) + progress.snapshot.StorePercentage(0) m.progress[runtimeID] = progress m.mu.Unlock() return token, ch, nil @@ -65,7 +65,7 @@ func (m *WorkerCallbackManager) Cancel(runtimeID string) { m.mu.Unlock() } -func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, percentage int64) error { +func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, percentage float64) error { m.mu.Lock() defer m.mu.Unlock() action, ok := m.actions[runtimeID] @@ -75,7 +75,9 @@ func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, p if token == "" || token != action.token { return fmt.Errorf("invalid worker token") } - m.progress[runtimeID].snapshot.Percentage.Store(max(0, min(100, percentage))) + if !m.progress[runtimeID].snapshot.StorePercentage(percentage) { + return fmt.Errorf("invalid percentage") + } return nil } @@ -86,7 +88,7 @@ func (m *WorkerCallbackManager) Progress(runtimeID string) (float64, bool) { if !ok { return 0, false } - return float64(progress.snapshot.Percentage.Load()), true + return progress.snapshot.Percentage(), true } func (m *WorkerCallbackManager) TrackSnapshotProgress(runtimeID string) *domain.SnapshotProgress { diff --git a/apps/druid/core/services/worker_callbacks_test.go b/apps/druid/core/services/worker_callbacks_test.go index 37157f99..069067c0 100644 --- a/apps/druid/core/services/worker_callbacks_test.go +++ b/apps/druid/core/services/worker_callbacks_test.go @@ -64,14 +64,14 @@ func TestWorkerCallbackTracksPullProgress(t *testing.T) { if progress, ok := manager.Progress("scroll-a"); !ok || progress != 0 { t.Fatalf("initial progress = %v, %v; want 0, true", progress, ok) } - if err := manager.ReportProgress("scroll-a", "wrong-token", 42); err == nil { + if err := manager.ReportProgress("scroll-a", "wrong-token", 42.3); err == nil { t.Fatal("invalid progress token should fail") } - if err := manager.ReportProgress("scroll-a", token, 42); err != nil { + if err := manager.ReportProgress("scroll-a", token, 42.3); err != nil { t.Fatal(err) } - if progress, ok := manager.Progress("scroll-a"); !ok || progress != 42 { - t.Fatalf("reported progress = %v, %v; want 42, true", progress, ok) + if progress, ok := manager.Progress("scroll-a"); !ok || progress != 42.3 { + t.Fatalf("reported progress = %v, %v; want 42.3, true", progress, ok) } manager.Cancel("scroll-a") @@ -86,10 +86,10 @@ func TestWorkerCallbackReadsTrackedSnapshotProgress(t *testing.T) { if again := manager.TrackSnapshotProgress("scroll-a"); again != progress { t.Fatal("tracking the same runtime replaced SnapshotProgress") } - progress.Percentage.Store(43) + progress.StorePercentage(43.5) - if got, ok := manager.Progress("scroll-a"); !ok || got != 43 { - t.Fatalf("progress = %v, %v; want 43, true", got, ok) + if got, ok := manager.Progress("scroll-a"); !ok || got != 43.5 { + t.Fatalf("progress = %v, %v; want 43.5, true", got, ok) } manager.ClearSnapshotProgress("scroll-a", progress) @@ -115,19 +115,19 @@ func TestWorkerCallbackKeepsTrackedSnapshotAcrossWorkerLifecycle(t *testing.T) { if registered != progress { t.Fatal("worker registration replaced tracked SnapshotProgress") } - if err := manager.ReportProgress("scroll-a", token, 42); err != nil { + if err := manager.ReportProgress("scroll-a", token, 42.3); err != nil { t.Fatal(err) } - if got := progress.Percentage.Load(); got != 42 { - t.Fatalf("tracked percentage = %d; want 42", got) + if got := progress.Percentage(); got != 42.3 { + t.Fatalf("tracked percentage = %v; want 42.3", got) } if err := manager.Complete("scroll-a", token, ports.RuntimeWorkerResult{}); err != nil { t.Fatal(err) } - progress.Percentage.Store(43) - if got, ok := manager.Progress("scroll-a"); !ok || got != 43 { - t.Fatalf("progress after worker completion = %v, %v; want 43, true", got, ok) + progress.StorePercentage(43.5) + if got, ok := manager.Progress("scroll-a"); !ok || got != 43.5 { + t.Fatalf("progress after worker completion = %v, %v; want 43.5, true", got, ok) } manager.ClearSnapshotProgress("scroll-a", progress) if _, ok := manager.Progress("scroll-a"); ok { diff --git a/internal/core/domain/oci.go b/internal/core/domain/oci.go index dfa400d2..07d1bacf 100644 --- a/internal/core/domain/oci.go +++ b/internal/core/domain/oci.go @@ -1,6 +1,9 @@ package domain -import "sync/atomic" +import ( + "math" + "sync/atomic" +) type ArtifactType string @@ -17,9 +20,11 @@ const ( SnapshotProgressModeRestore = "restore" ) +const snapshotProgressScale int64 = 10 + // SnapshotProgress tracks the state of a data pull/push operation. type SnapshotProgress struct { - Percentage atomic.Int64 + percentage atomic.Int64 Mode atomic.Value // stores string } @@ -29,6 +34,19 @@ func NewSnapshotProgress() *SnapshotProgress { return sp } +func (p *SnapshotProgress) StorePercentage(value float64) bool { + if math.IsNaN(value) || math.IsInf(value, 0) { + return false + } + value = math.Max(0, math.Min(100, value)) + p.percentage.Store(int64(math.Round(value * float64(snapshotProgressScale)))) + return true +} + +func (p *SnapshotProgress) Percentage() float64 { + return float64(p.percentage.Load()) / float64(snapshotProgressScale) +} + type AnnotationInfo struct { MinRam string MinDisk string diff --git a/internal/core/domain/oci_test.go b/internal/core/domain/oci_test.go new file mode 100644 index 00000000..f6f9645a --- /dev/null +++ b/internal/core/domain/oci_test.go @@ -0,0 +1,31 @@ +package domain + +import ( + "math" + "testing" +) + +func TestSnapshotProgressStoresTenths(t *testing.T) { + progress := NewSnapshotProgress() + if !progress.StorePercentage(54.71) { + t.Fatal("valid percentage was rejected") + } + if got := progress.Percentage(); got != 54.7 { + t.Fatalf("percentage = %v; want 54.7", got) + } +} + +func TestSnapshotProgressClampsAndRejectsInvalidValues(t *testing.T) { + progress := NewSnapshotProgress() + progress.StorePercentage(150) + if got := progress.Percentage(); got != 100 { + t.Fatalf("percentage = %v; want 100", got) + } + progress.StorePercentage(-10) + if got := progress.Percentage(); got != 0 { + t.Fatalf("percentage = %v; want 0", got) + } + if progress.StorePercentage(math.NaN()) { + t.Fatal("NaN was accepted") + } +} diff --git a/internal/core/services/coldstarter/handler/lua_handler.go b/internal/core/services/coldstarter/handler/lua_handler.go index cdbda1a7..c3d49b1e 100644 --- a/internal/core/services/coldstarter/handler/lua_handler.go +++ b/internal/core/services/coldstarter/handler/lua_handler.go @@ -188,7 +188,7 @@ func (handler *LuaHandler) GetHandler(funcs map[string]func(data ...string)) (po l.SetGlobal("get_snapshot_percentage", l.NewFunction( func(l *lua.LState) int { if handler.progress != nil { - l.Push(lua.LNumber(handler.progress.Percentage.Load())) + l.Push(lua.LNumber(handler.progress.Percentage())) } else { l.Push(lua.LNumber(100)) } diff --git a/internal/core/services/registry/oci.go b/internal/core/services/registry/oci.go index 226ebfda..cdb2a9aa 100644 --- a/internal/core/services/registry/oci.go +++ b/internal/core/services/registry/oci.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "math" "net/http" "os" "path" @@ -200,14 +201,14 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, if progress != nil { progress.Mode.Store(domain.SnapshotProgressModeRestore) - progress.Percentage.Store(0) + progress.StorePercentage(0) defer progress.Mode.Store(domain.SnapshotProgressModeIdle) } storeProgress := func(done, total int64) { if progress == nil || total <= 0 { return } - progress.Percentage.Store(min(99, done*100/total)) + progress.StorePercentage(math.Min(99, float64(done)*100/float64(total))) } copyOpts := oras.CopyOptions{ @@ -343,7 +344,7 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, } if progress != nil { - progress.Percentage.Store(100) + progress.StorePercentage(100) } return nil diff --git a/internal/core/services/registry/oci_test.go b/internal/core/services/registry/oci_test.go index e8d5a841..ee1f137c 100644 --- a/internal/core/services/registry/oci_test.go +++ b/internal/core/services/registry/oci_test.go @@ -150,8 +150,8 @@ func TestPullSelectiveDoesNotReportCompleteBeforeMetadataIsWritten(t *testing.T) if err := client.PullSelective(destination, repository+":1.0", true, progress); err == nil { t.Fatal("pull should fail when manifest.json cannot be written") } - if percentage := progress.Percentage.Load(); percentage >= 100 { - t.Fatalf("failed pull progress = %d; want below 100", percentage) + if percentage := progress.Percentage(); percentage >= 100 { + t.Fatalf("failed pull progress = %v; want below 100", percentage) } if mode := progress.Mode.Load(); mode != domain.SnapshotProgressModeIdle { t.Fatalf("failed pull mode = %v; want idle", mode) diff --git a/internal/runtime/kubernetes/snapshot_progress.go b/internal/runtime/kubernetes/snapshot_progress.go index 7b70a985..ff37f224 100644 --- a/internal/runtime/kubernetes/snapshot_progress.go +++ b/internal/runtime/kubernetes/snapshot_progress.go @@ -55,7 +55,5 @@ func storeSnapshotProgress(progress *domain.SnapshotProgress, current float64, t math.IsNaN(total) || math.IsInf(total, 0) { return false } - percentage := int64(math.Round(current / total * 100)) - progress.Percentage.Store(max(0, min(100, percentage))) - return true + return progress.StorePercentage(current / total * 100) } diff --git a/internal/runtime/kubernetes/snapshot_progress_test.go b/internal/runtime/kubernetes/snapshot_progress_test.go index ff4125d3..941adda3 100644 --- a/internal/runtime/kubernetes/snapshot_progress_test.go +++ b/internal/runtime/kubernetes/snapshot_progress_test.go @@ -15,8 +15,20 @@ func TestObserveSnapshotProgressReadsSteamCMDByteMarker(t *testing.T) { if !observeSnapshotProgress(line, progress) { t.Fatal("progress marker was not recognized") } - if got := progress.Percentage.Load(); got != 55 { - t.Fatalf("percentage = %d; want 55", got) + if got := progress.Percentage(); got != 54.7 { + t.Fatalf("percentage = %v; want 54.7", got) + } +} + +func TestObserveSnapshotProgressReadsLowProgressByteMarker(t *testing.T) { + progress := domain.NewSnapshotProgress() + line := `DRUID_PROGRESS_V1 {"unit":"bytes","current":401991074,"total":22938933947}` + + if !observeSnapshotProgress(line, progress) { + t.Fatal("progress marker was not recognized") + } + if got := progress.Percentage(); got != 1.8 { + t.Fatalf("percentage = %v; want 1.8", got) } } @@ -27,8 +39,8 @@ func TestObserveSnapshotProgressReadsOriginalSteamCMDLine(t *testing.T) { if observeSnapshotProgress(line, progress) { t.Fatal("original SteamCMD output should remain visible in the console") } - if got := progress.Percentage.Load(); got != 55 { - t.Fatalf("percentage = %d; want 55", got) + if got := progress.Percentage(); got != 54.7 { + t.Fatalf("percentage = %v; want 54.7", got) } } @@ -54,8 +66,8 @@ func TestReadSnapshotProgressDoesNotDependOnConsoleConsumer(t *testing.T) { case <-time.After(100 * time.Millisecond): t.Fatal("progress processing blocked behind the console backlog") } - if got := progress.Percentage.Load(); got != 55 { - t.Fatalf("percentage = %d; want 55", got) + if got := progress.Percentage(); got != 54.7 { + t.Fatalf("percentage = %v; want 54.7", got) } if got := <-output; got != "older console line" { t.Fatalf("console backlog changed to %q", got) @@ -64,7 +76,7 @@ func TestReadSnapshotProgressDoesNotDependOnConsoleConsumer(t *testing.T) { func TestObserveSnapshotProgressLeavesMalformedLinesAlone(t *testing.T) { progress := domain.NewSnapshotProgress() - progress.Percentage.Store(17) + progress.StorePercentage(17) for _, line := range []string{ "ordinary server output", @@ -74,8 +86,8 @@ func TestObserveSnapshotProgressLeavesMalformedLinesAlone(t *testing.T) { if observeSnapshotProgress(line, progress) { t.Fatalf("line was unexpectedly recognized: %s", line) } - if got := progress.Percentage.Load(); got != 17 { - t.Fatalf("percentage changed to %d for line %q", got, line) + if got := progress.Percentage(); got != 17 { + t.Fatalf("percentage changed to %v for line %q", got, line) } } } From da0568db2ffaf7c3e88658c406110bd197e6682b Mon Sep 17 00:00:00 2001 From: IOHelpMe <101821466+IOHelpMe@users.noreply.github.com> Date: Thu, 30 Jul 2026 21:18:22 +0200 Subject: [PATCH 4/4] fix(runtime): rely only on SnapshotProgress --- apps/druid/adapters/cli/callback.go | 11 +-- apps/druid/adapters/cli/callback_test.go | 6 +- .../adapters/cli/worker_progress_test.go | 12 +-- apps/druid/adapters/cli/worker_pull.go | 8 +- apps/druid/core/services/runtime_lifecycle.go | 6 -- apps/druid/core/services/runtime_session.go | 1 - .../core/services/runtime_session_cache.go | 6 -- .../services/runtime_session_execution.go | 15 ++- apps/druid/core/services/runtime_update.go | 3 - apps/druid/core/services/worker_callbacks.go | 64 ++----------- .../core/services/worker_callbacks_test.go | 63 +------------ internal/core/domain/oci.go | 22 +---- internal/core/domain/oci_test.go | 31 ------- internal/core/ports/services_ports.go | 1 - .../coldstarter/handler/lua_handler.go | 2 +- internal/core/services/registry/oci.go | 7 +- internal/core/services/registry/oci_test.go | 4 +- internal/runtime/kubernetes/procedures.go | 8 +- .../runtime/kubernetes/snapshot_progress.go | 59 ------------ .../kubernetes/snapshot_progress_test.go | 93 ------------------- internal/runtime/kubernetes/wait_jobs.go | 75 ++++----------- 21 files changed, 66 insertions(+), 431 deletions(-) delete mode 100644 internal/core/domain/oci_test.go delete mode 100644 internal/runtime/kubernetes/snapshot_progress.go delete mode 100644 internal/runtime/kubernetes/snapshot_progress_test.go diff --git a/apps/druid/adapters/cli/callback.go b/apps/druid/adapters/cli/callback.go index 8a9ab3af..269cd4a0 100644 --- a/apps/druid/adapters/cli/callback.go +++ b/apps/druid/adapters/cli/callback.go @@ -1,8 +1,6 @@ package cli import ( - "math" - "github.com/gofiber/fiber/v2" appservices "github.com/highcard-dev/daemon/apps/druid/core/services" "github.com/highcard-dev/daemon/internal/callbackapi" @@ -15,17 +13,16 @@ type runtimeCallbackHandler struct { func (h runtimeCallbackHandler) ReportProgress(c *fiber.Ctx) error { var report struct { - Token string `json:"token"` - Percentage *float64 `json:"percentage"` + Token string `json:"token"` + Percentage *int64 `json:"percentage"` } if err := c.BodyParser(&report); err != nil || report.Percentage == nil { return fiber.NewError(fiber.StatusBadRequest, "invalid progress report") } - percentage := *report.Percentage - if math.IsNaN(percentage) || math.IsInf(percentage, 0) || percentage < 0 || percentage > 100 { + if *report.Percentage < 0 || *report.Percentage > 100 { return fiber.NewError(fiber.StatusBadRequest, "percentage must be between 0 and 100") } - if err := h.callbacks.ReportProgress(c.Params("runtime_id"), report.Token, percentage); err != nil { + if err := h.callbacks.ReportProgress(c.Params("runtime_id"), report.Token, *report.Percentage); err != nil { return fiber.NewError(fiber.StatusUnauthorized, err.Error()) } return c.SendStatus(fiber.StatusNoContent) diff --git a/apps/druid/adapters/cli/callback_test.go b/apps/druid/adapters/cli/callback_test.go index 0b9465bb..259e11b6 100644 --- a/apps/druid/adapters/cli/callback_test.go +++ b/apps/druid/adapters/cli/callback_test.go @@ -24,7 +24,7 @@ func TestRuntimeCallbackHandlerReportsProgress(t *testing.T) { request := httptest.NewRequest( http.MethodPost, "/internal/v1/workers/runtime-1/progress", - strings.NewReader(fmt.Sprintf(`{"token":%q,"percentage":42.3}`, token)), + strings.NewReader(fmt.Sprintf(`{"token":%q,"percentage":42}`, token)), ) request.Header.Set(fiber.HeaderContentType, fiber.MIMEApplicationJSON) response, err := app.Test(request) @@ -34,7 +34,7 @@ func TestRuntimeCallbackHandlerReportsProgress(t *testing.T) { if response.StatusCode != http.StatusNoContent { t.Fatalf("status = %d; want %d", response.StatusCode, http.StatusNoContent) } - if progress, ok := callbacks.Progress("runtime-1"); !ok || progress != 42.3 { - t.Fatalf("progress = %v, %v; want 42.3, true", progress, ok) + if progress, ok := callbacks.Progress("runtime-1"); !ok || progress != 42 { + t.Fatalf("progress = %v, %v; want 42, true", progress, ok) } } diff --git a/apps/druid/adapters/cli/worker_progress_test.go b/apps/druid/adapters/cli/worker_progress_test.go index c5891e92..42f7818c 100644 --- a/apps/druid/adapters/cli/worker_progress_test.go +++ b/apps/druid/adapters/cli/worker_progress_test.go @@ -12,11 +12,11 @@ import ( ) func TestWorkerProgressReporterReadsSnapshotProgress(t *testing.T) { - reports := make(chan float64, 1) + reports := make(chan int64, 1) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { var report struct { - Token string `json:"token"` - Percentage float64 `json:"percentage"` + Token string `json:"token"` + Percentage int64 `json:"percentage"` } if err := json.NewDecoder(request.Body).Decode(&report); err != nil { t.Error(err) @@ -30,7 +30,7 @@ func TestWorkerProgressReporterReadsSnapshotProgress(t *testing.T) { defer server.Close() progress := domain.NewSnapshotProgress() - progress.StorePercentage(37.4) + progress.Percentage.Store(37) stop := startWorkerProgressReporter( ports.RuntimeWorkerAction{ RuntimeID: "runtime-1", @@ -44,8 +44,8 @@ func TestWorkerProgressReporterReadsSnapshotProgress(t *testing.T) { select { case percentage := <-reports: - if percentage != 37.4 { - t.Fatalf("percentage = %v; want 37.4", percentage) + if percentage != 37 { + t.Fatalf("percentage = %d; want 37", percentage) } case <-time.After(time.Second): t.Fatal("progress was not reported") diff --git a/apps/druid/adapters/cli/worker_pull.go b/apps/druid/adapters/cli/worker_pull.go index dd0da557..1a9ec26f 100644 --- a/apps/druid/adapters/cli/worker_pull.go +++ b/apps/druid/adapters/cli/worker_pull.go @@ -338,15 +338,15 @@ func startWorkerProgressReporter(action ports.RuntimeWorkerAction, progress *dom defer wait.Done() ticker := time.NewTicker(interval) defer ticker.Stop() - lastPercentage := -1.0 + lastPercentage := int64(-1) report := func() { - percentage := progress.Percentage() + percentage := progress.Percentage.Load() if percentage == lastPercentage { return } body, _ := json.Marshal(struct { - Token string `json:"token"` - Percentage float64 `json:"percentage"` + Token string `json:"token"` + Percentage int64 `json:"percentage"` }{action.CallbackToken, percentage}) request, _ := http.NewRequest(http.MethodPost, progressURL, bytes.NewReader(body)) request.Header.Set("Content-Type", "application/json") diff --git a/apps/druid/core/services/runtime_lifecycle.go b/apps/druid/core/services/runtime_lifecycle.go index 9c7730d4..a0df0ef3 100644 --- a/apps/druid/core/services/runtime_lifecycle.go +++ b/apps/druid/core/services/runtime_lifecycle.go @@ -13,9 +13,6 @@ func (s *RuntimeSupervisor) DeleteWithPolicy(id string, purgeData bool) error { s.mu.Unlock() if session != nil { session.stopDeploymentQueue() - if s.workerCallbacks != nil { - s.workerCallbacks.ClearSnapshotProgress(id, session.snapshotProgress) - } } runtimeScroll, err := s.store.GetScroll(id) @@ -63,9 +60,6 @@ func (s *RuntimeSupervisor) Stop(id string) (*domain.RuntimeScroll, error) { session.markError(err) return nil, err } - if s.workerCallbacks != nil { - s.workerCallbacks.ClearSnapshotProgress(id, session.snapshotProgress) - } session.stopDeploymentQueue() return s.store.GetScroll(id) } diff --git a/apps/druid/core/services/runtime_session.go b/apps/druid/core/services/runtime_session.go index 4c140129..83e31d25 100644 --- a/apps/druid/core/services/runtime_session.go +++ b/apps/druid/core/services/runtime_session.go @@ -19,7 +19,6 @@ type RuntimeSession struct { scrollService *coreservices.ScrollService watchService ports.WatchServiceInterface runtimeBackend ports.RuntimeBackendInterface - snapshotProgress *domain.SnapshotProgress queue map[string]*runtimeQueueItem workWg sync.WaitGroup notifierChan []chan []string diff --git a/apps/druid/core/services/runtime_session_cache.go b/apps/druid/core/services/runtime_session_cache.go index 669586e4..8b4063e2 100644 --- a/apps/druid/core/services/runtime_session_cache.go +++ b/apps/druid/core/services/runtime_session_cache.go @@ -54,18 +54,12 @@ func (s *RuntimeSupervisor) startSession(runtimeScroll *domain.RuntimeScroll) (* session.devDaemonToken = s.internalToken session.devAuthJWKSURL = s.authJWKSURL session.devRuntimeJWKSURL = s.runtimeJWKSURL - if s.workerCallbacks != nil { - session.snapshotProgress = s.workerCallbacks.TrackSnapshotProgress(runtimeScroll.ID) - } session.Start() s.mu.Lock() if existing := s.sessions[runtimeScroll.ID]; existing != nil { s.mu.Unlock() session.stopDeploymentQueue() - if s.workerCallbacks != nil { - s.workerCallbacks.ClearSnapshotProgress(runtimeScroll.ID, session.snapshotProgress) - } return existing, nil } s.sessions[runtimeScroll.ID] = session diff --git a/apps/druid/core/services/runtime_session_execution.go b/apps/druid/core/services/runtime_session_execution.go index 5c7961a5..b845c91c 100644 --- a/apps/druid/core/services/runtime_session_execution.go +++ b/apps/druid/core/services/runtime_session_execution.go @@ -52,14 +52,13 @@ func (s *RuntimeSession) runCommand(cmd string) error { } exitCode, err := s.runtimeBackend.RunCommand(ports.RuntimeCommand{ - Name: cmd, - ScrollID: scrollID, - Command: command, - Root: root, - GlobalPorts: runtimePorts, - Routing: routing, - ProcedureEnv: procedureEnv, - SnapshotProgress: s.snapshotProgress, + Name: cmd, + ScrollID: scrollID, + Command: command, + Root: root, + GlobalPorts: runtimePorts, + Routing: routing, + ProcedureEnv: procedureEnv, ProcedureStatusObserver: func(procedure string, status domain.ScrollLockStatus, exitCode *int) { s.persistProcedureStatus(cmd, procedure, status, exitCode) }, diff --git a/apps/druid/core/services/runtime_update.go b/apps/druid/core/services/runtime_update.go index 92190953..03e22bf9 100644 --- a/apps/druid/core/services/runtime_update.go +++ b/apps/druid/core/services/runtime_update.go @@ -33,9 +33,6 @@ func (s *RuntimeSupervisor) updateExistingScroll(runtimeScroll *domain.RuntimeSc s.mu.Unlock() if session != nil { session.stopDeploymentQueue() - if s.workerCallbacks != nil { - s.workerCallbacks.ClearSnapshotProgress(runtimeScroll.ID, session.snapshotProgress) - } } if wasRunning { diff --git a/apps/druid/core/services/worker_callbacks.go b/apps/druid/core/services/worker_callbacks.go index 470c6152..dba37751 100644 --- a/apps/druid/core/services/worker_callbacks.go +++ b/apps/druid/core/services/worker_callbacks.go @@ -6,14 +6,13 @@ import ( "fmt" "sync" - "github.com/highcard-dev/daemon/internal/core/domain" "github.com/highcard-dev/daemon/internal/core/ports" ) type WorkerCallbackManager struct { mu sync.Mutex actions map[string]workerCallbackAction - progress map[string]workerCallbackProgress + progress map[string]int64 } type workerCallbackAction struct { @@ -21,15 +20,10 @@ type workerCallbackAction struct { result chan ports.RuntimeWorkerResult } -type workerCallbackProgress struct { - snapshot *domain.SnapshotProgress - trackers int -} - func NewWorkerCallbackManager() *WorkerCallbackManager { return &WorkerCallbackManager{ actions: map[string]workerCallbackAction{}, - progress: map[string]workerCallbackProgress{}, + progress: map[string]int64{}, } } @@ -46,12 +40,7 @@ func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports return "", nil, fmt.Errorf("worker action already pending for runtime %s", runtimeID) } m.actions[runtimeID] = workerCallbackAction{token: token, result: ch} - progress := m.progress[runtimeID] - if progress.snapshot == nil { - progress.snapshot = domain.NewSnapshotProgress() - } - progress.snapshot.StorePercentage(0) - m.progress[runtimeID] = progress + m.progress[runtimeID] = 0 m.mu.Unlock() return token, ch, nil } @@ -59,13 +48,11 @@ func (m *WorkerCallbackManager) Register(runtimeID string) (string, <-chan ports func (m *WorkerCallbackManager) Cancel(runtimeID string) { m.mu.Lock() delete(m.actions, runtimeID) - if progress := m.progress[runtimeID]; progress.trackers == 0 { - delete(m.progress, runtimeID) - } + delete(m.progress, runtimeID) m.mu.Unlock() } -func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, percentage float64) error { +func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, percentage int64) error { m.mu.Lock() defer m.mu.Unlock() action, ok := m.actions[runtimeID] @@ -75,9 +62,7 @@ func (m *WorkerCallbackManager) ReportProgress(runtimeID string, token string, p if token == "" || token != action.token { return fmt.Errorf("invalid worker token") } - if !m.progress[runtimeID].snapshot.StorePercentage(percentage) { - return fmt.Errorf("invalid percentage") - } + m.progress[runtimeID] = max(0, min(100, percentage)) return nil } @@ -85,38 +70,7 @@ func (m *WorkerCallbackManager) Progress(runtimeID string) (float64, bool) { m.mu.Lock() defer m.mu.Unlock() progress, ok := m.progress[runtimeID] - if !ok { - return 0, false - } - return progress.snapshot.Percentage(), true -} - -func (m *WorkerCallbackManager) TrackSnapshotProgress(runtimeID string) *domain.SnapshotProgress { - m.mu.Lock() - progress := m.progress[runtimeID] - if progress.snapshot == nil { - progress.snapshot = domain.NewSnapshotProgress() - } - progress.trackers++ - m.progress[runtimeID] = progress - m.mu.Unlock() - return progress.snapshot -} - -func (m *WorkerCallbackManager) ClearSnapshotProgress(runtimeID string, progress *domain.SnapshotProgress) { - m.mu.Lock() - current, ok := m.progress[runtimeID] - if ok && current.snapshot == progress { - if current.trackers > 0 { - current.trackers-- - } - if _, pending := m.actions[runtimeID]; pending || current.trackers > 0 { - m.progress[runtimeID] = current - } else { - delete(m.progress, runtimeID) - } - } - m.mu.Unlock() + return float64(progress), ok } func (m *WorkerCallbackManager) Complete(runtimeID string, token string, result ports.RuntimeWorkerResult) error { @@ -131,9 +85,7 @@ func (m *WorkerCallbackManager) Complete(runtimeID string, token string, result return fmt.Errorf("invalid worker token") } delete(m.actions, runtimeID) - if progress := m.progress[runtimeID]; progress.trackers == 0 { - delete(m.progress, runtimeID) - } + delete(m.progress, runtimeID) m.mu.Unlock() action.result <- result close(action.result) diff --git a/apps/druid/core/services/worker_callbacks_test.go b/apps/druid/core/services/worker_callbacks_test.go index 069067c0..c9dd517e 100644 --- a/apps/druid/core/services/worker_callbacks_test.go +++ b/apps/druid/core/services/worker_callbacks_test.go @@ -64,14 +64,14 @@ func TestWorkerCallbackTracksPullProgress(t *testing.T) { if progress, ok := manager.Progress("scroll-a"); !ok || progress != 0 { t.Fatalf("initial progress = %v, %v; want 0, true", progress, ok) } - if err := manager.ReportProgress("scroll-a", "wrong-token", 42.3); err == nil { + if err := manager.ReportProgress("scroll-a", "wrong-token", 42); err == nil { t.Fatal("invalid progress token should fail") } - if err := manager.ReportProgress("scroll-a", token, 42.3); err != nil { + if err := manager.ReportProgress("scroll-a", token, 42); err != nil { t.Fatal(err) } - if progress, ok := manager.Progress("scroll-a"); !ok || progress != 42.3 { - t.Fatalf("reported progress = %v, %v; want 42.3, true", progress, ok) + if progress, ok := manager.Progress("scroll-a"); !ok || progress != 42 { + t.Fatalf("reported progress = %v, %v; want 42, true", progress, ok) } manager.Cancel("scroll-a") @@ -79,58 +79,3 @@ func TestWorkerCallbackTracksPullProgress(t *testing.T) { t.Fatal("cancelled progress should be removed") } } - -func TestWorkerCallbackReadsTrackedSnapshotProgress(t *testing.T) { - manager := NewWorkerCallbackManager() - progress := manager.TrackSnapshotProgress("scroll-a") - if again := manager.TrackSnapshotProgress("scroll-a"); again != progress { - t.Fatal("tracking the same runtime replaced SnapshotProgress") - } - progress.StorePercentage(43.5) - - if got, ok := manager.Progress("scroll-a"); !ok || got != 43.5 { - t.Fatalf("progress = %v, %v; want 43.5, true", got, ok) - } - - manager.ClearSnapshotProgress("scroll-a", progress) - if _, ok := manager.Progress("scroll-a"); !ok { - t.Fatal("clearing one tracker removed another tracker's SnapshotProgress") - } - manager.ClearSnapshotProgress("scroll-a", progress) - if _, ok := manager.Progress("scroll-a"); ok { - t.Fatal("cleared SnapshotProgress should be removed") - } -} - -func TestWorkerCallbackKeepsTrackedSnapshotAcrossWorkerLifecycle(t *testing.T) { - manager := NewWorkerCallbackManager() - progress := manager.TrackSnapshotProgress("scroll-a") - token, _, err := manager.Register("scroll-a") - if err != nil { - t.Fatal(err) - } - manager.mu.Lock() - registered := manager.progress["scroll-a"].snapshot - manager.mu.Unlock() - if registered != progress { - t.Fatal("worker registration replaced tracked SnapshotProgress") - } - if err := manager.ReportProgress("scroll-a", token, 42.3); err != nil { - t.Fatal(err) - } - if got := progress.Percentage(); got != 42.3 { - t.Fatalf("tracked percentage = %v; want 42.3", got) - } - if err := manager.Complete("scroll-a", token, ports.RuntimeWorkerResult{}); err != nil { - t.Fatal(err) - } - - progress.StorePercentage(43.5) - if got, ok := manager.Progress("scroll-a"); !ok || got != 43.5 { - t.Fatalf("progress after worker completion = %v, %v; want 43.5, true", got, ok) - } - manager.ClearSnapshotProgress("scroll-a", progress) - if _, ok := manager.Progress("scroll-a"); ok { - t.Fatal("cleared tracked progress should be removed") - } -} diff --git a/internal/core/domain/oci.go b/internal/core/domain/oci.go index 07d1bacf..dfa400d2 100644 --- a/internal/core/domain/oci.go +++ b/internal/core/domain/oci.go @@ -1,9 +1,6 @@ package domain -import ( - "math" - "sync/atomic" -) +import "sync/atomic" type ArtifactType string @@ -20,11 +17,9 @@ const ( SnapshotProgressModeRestore = "restore" ) -const snapshotProgressScale int64 = 10 - // SnapshotProgress tracks the state of a data pull/push operation. type SnapshotProgress struct { - percentage atomic.Int64 + Percentage atomic.Int64 Mode atomic.Value // stores string } @@ -34,19 +29,6 @@ func NewSnapshotProgress() *SnapshotProgress { return sp } -func (p *SnapshotProgress) StorePercentage(value float64) bool { - if math.IsNaN(value) || math.IsInf(value, 0) { - return false - } - value = math.Max(0, math.Min(100, value)) - p.percentage.Store(int64(math.Round(value * float64(snapshotProgressScale)))) - return true -} - -func (p *SnapshotProgress) Percentage() float64 { - return float64(p.percentage.Load()) / float64(snapshotProgressScale) -} - type AnnotationInfo struct { MinRam string MinDisk string diff --git a/internal/core/domain/oci_test.go b/internal/core/domain/oci_test.go deleted file mode 100644 index f6f9645a..00000000 --- a/internal/core/domain/oci_test.go +++ /dev/null @@ -1,31 +0,0 @@ -package domain - -import ( - "math" - "testing" -) - -func TestSnapshotProgressStoresTenths(t *testing.T) { - progress := NewSnapshotProgress() - if !progress.StorePercentage(54.71) { - t.Fatal("valid percentage was rejected") - } - if got := progress.Percentage(); got != 54.7 { - t.Fatalf("percentage = %v; want 54.7", got) - } -} - -func TestSnapshotProgressClampsAndRejectsInvalidValues(t *testing.T) { - progress := NewSnapshotProgress() - progress.StorePercentage(150) - if got := progress.Percentage(); got != 100 { - t.Fatalf("percentage = %v; want 100", got) - } - progress.StorePercentage(-10) - if got := progress.Percentage(); got != 0 { - t.Fatalf("percentage = %v; want 0", got) - } - if progress.StorePercentage(math.NaN()) { - t.Fatal("NaN was accepted") - } -} diff --git a/internal/core/ports/services_ports.go b/internal/core/ports/services_ports.go index cd0a642f..5de1eb3f 100644 --- a/internal/core/ports/services_ports.go +++ b/internal/core/ports/services_ports.go @@ -80,7 +80,6 @@ type RuntimeCommand struct { GlobalPorts []domain.Port Routing []domain.RuntimeRouteAssignment ProcedureEnv map[string]map[string]string - SnapshotProgress *domain.SnapshotProgress ProcedureStatusObserver func(procedure string, status domain.ScrollLockStatus, exitCode *int) } diff --git a/internal/core/services/coldstarter/handler/lua_handler.go b/internal/core/services/coldstarter/handler/lua_handler.go index c3d49b1e..cdbda1a7 100644 --- a/internal/core/services/coldstarter/handler/lua_handler.go +++ b/internal/core/services/coldstarter/handler/lua_handler.go @@ -188,7 +188,7 @@ func (handler *LuaHandler) GetHandler(funcs map[string]func(data ...string)) (po l.SetGlobal("get_snapshot_percentage", l.NewFunction( func(l *lua.LState) int { if handler.progress != nil { - l.Push(lua.LNumber(handler.progress.Percentage())) + l.Push(lua.LNumber(handler.progress.Percentage.Load())) } else { l.Push(lua.LNumber(100)) } diff --git a/internal/core/services/registry/oci.go b/internal/core/services/registry/oci.go index cdb2a9aa..226ebfda 100644 --- a/internal/core/services/registry/oci.go +++ b/internal/core/services/registry/oci.go @@ -5,7 +5,6 @@ import ( "encoding/json" "errors" "fmt" - "math" "net/http" "os" "path" @@ -201,14 +200,14 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, if progress != nil { progress.Mode.Store(domain.SnapshotProgressModeRestore) - progress.StorePercentage(0) + progress.Percentage.Store(0) defer progress.Mode.Store(domain.SnapshotProgressModeIdle) } storeProgress := func(done, total int64) { if progress == nil || total <= 0 { return } - progress.StorePercentage(math.Min(99, float64(done)*100/float64(total))) + progress.Percentage.Store(min(99, done*100/total)) } copyOpts := oras.CopyOptions{ @@ -344,7 +343,7 @@ func (c *OciClient) PullSelective(dir string, artifact string, includeData bool, } if progress != nil { - progress.StorePercentage(100) + progress.Percentage.Store(100) } return nil diff --git a/internal/core/services/registry/oci_test.go b/internal/core/services/registry/oci_test.go index ee1f137c..e8d5a841 100644 --- a/internal/core/services/registry/oci_test.go +++ b/internal/core/services/registry/oci_test.go @@ -150,8 +150,8 @@ func TestPullSelectiveDoesNotReportCompleteBeforeMetadataIsWritten(t *testing.T) if err := client.PullSelective(destination, repository+":1.0", true, progress); err == nil { t.Fatal("pull should fail when manifest.json cannot be written") } - if percentage := progress.Percentage(); percentage >= 100 { - t.Fatalf("failed pull progress = %v; want below 100", percentage) + if percentage := progress.Percentage.Load(); percentage >= 100 { + t.Fatalf("failed pull progress = %d; want below 100", percentage) } if mode := progress.Mode.Load(); mode != domain.SnapshotProgressModeIdle { t.Fatalf("failed pull mode = %v; want idle", mode) diff --git a/internal/runtime/kubernetes/procedures.go b/internal/runtime/kubernetes/procedures.go index 0cb745d7..6b02c2d9 100644 --- a/internal/runtime/kubernetes/procedures.go +++ b/internal/runtime/kubernetes/procedures.go @@ -108,7 +108,7 @@ func (b *Backend) RunCommand(command ports.RuntimeCommand) (*int, error) { continue } command.ObserveProcedureStatus(procedureName, domain.ScrollLockStatusRunning, nil) - exitCode, err := b.runJobProcedure(command.ScrollID, command.Name, procedureName, resourceName, procedure, command.Root, command.GlobalPorts, env, portUse, command.SnapshotProgress) + exitCode, err := b.runJobProcedure(command.ScrollID, command.Name, procedureName, resourceName, procedure, command.Root, command.GlobalPorts, env, portUse) if err != nil { if exitCode != nil && *exitCode != 0 && procedure.IgnoreFailure { command.ObserveProcedureStatus(procedureName, domain.ScrollLockStatusDone, exitCode) @@ -138,7 +138,7 @@ func (b *Backend) RunCommand(command ports.RuntimeCommand) (*int, error) { return nil, nil } -func (b *Backend) runJobProcedure(scrollID string, commandName string, procedureName string, resourceName string, procedure *domain.Procedure, root string, globalPorts []domain.Port, env map[string]string, portUse map[string]int, progress *domain.SnapshotProgress) (*int, error) { +func (b *Backend) runJobProcedure(scrollID string, commandName string, procedureName string, resourceName string, procedure *domain.Procedure, root string, globalPorts []domain.Port, env map[string]string, portUse map[string]int) (*int, error) { if procedure.IsSignal() { logger.Log().Info("Running Kubernetes signal procedure", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("target", procedure.Target), zap.String("signal", procedure.Signal)) if err := b.Signal(procedureName, procedure.Target, procedure.Signal, root); err != nil { @@ -196,7 +196,7 @@ func (b *Backend) runJobProcedure(scrollID string, commandName string, procedure if err == nil { streamStarted = true logger.Log().Debug("Streaming Kubernetes job procedure logs", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("namespace", namespace), zap.String("job", jobName), zap.String("pod", podName), zap.String("console_id", consoleID)) - go b.streamPodLogs(ctx, namespace, podName, output, progress) + go b.streamPodLogs(ctx, namespace, podName, output) } else { logger.Log().Warn("Could not find Kubernetes job pod before wait; console logs may be empty", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("namespace", namespace), zap.String("job", jobName), zap.Error(err)) } @@ -301,7 +301,7 @@ func (b *Backend) ensurePersistentProcedure(ctx context.Context, scrollID string return } logger.Log().Debug("Streaming Kubernetes persistent procedure logs", zap.String("scroll_id", scrollID), zap.String("command", commandName), zap.String("procedure", procedureName), zap.String("namespace", namespace), zap.String("pod", podName)) - b.streamPodLogs(context.Background(), namespace, podName, output, nil) + b.streamPodLogs(context.Background(), namespace, podName, output) }() return nil } diff --git a/internal/runtime/kubernetes/snapshot_progress.go b/internal/runtime/kubernetes/snapshot_progress.go deleted file mode 100644 index ff37f224..00000000 --- a/internal/runtime/kubernetes/snapshot_progress.go +++ /dev/null @@ -1,59 +0,0 @@ -package kubernetes - -import ( - "encoding/json" - "math" - "regexp" - "strconv" - "strings" - - "github.com/highcard-dev/daemon/internal/core/domain" -) - -const snapshotProgressPrefix = "DRUID_PROGRESS_V1 " - -var steamCMDProgressPattern = regexp.MustCompile( - `Update state \(0x[0-9a-fA-F]+\) downloading, progress: [0-9]+(?:\.[0-9]+)? \(([0-9]+) / ([0-9]+)\)`, -) - -type snapshotProgressSample struct { - Unit string `json:"unit"` - Current float64 `json:"current"` - Total float64 `json:"total"` -} - -func observeSnapshotProgress(line string, progress *domain.SnapshotProgress) bool { - if progress == nil { - return false - } - - if strings.HasPrefix(line, snapshotProgressPrefix) { - var sample snapshotProgressSample - if err := json.Unmarshal([]byte(strings.TrimPrefix(line, snapshotProgressPrefix)), &sample); err != nil { - return false - } - if sample.Unit != "bytes" || !storeSnapshotProgress(progress, sample.Current, sample.Total) { - return false - } - return true - } - - matches := steamCMDProgressPattern.FindStringSubmatch(line) - if len(matches) == 3 { - current, currentErr := strconv.ParseFloat(matches[1], 64) - total, totalErr := strconv.ParseFloat(matches[2], 64) - if currentErr == nil && totalErr == nil { - storeSnapshotProgress(progress, current, total) - } - } - return false -} - -func storeSnapshotProgress(progress *domain.SnapshotProgress, current float64, total float64) bool { - if current < 0 || total <= 0 || - math.IsNaN(current) || math.IsInf(current, 0) || - math.IsNaN(total) || math.IsInf(total, 0) { - return false - } - return progress.StorePercentage(current / total * 100) -} diff --git a/internal/runtime/kubernetes/snapshot_progress_test.go b/internal/runtime/kubernetes/snapshot_progress_test.go deleted file mode 100644 index 941adda3..00000000 --- a/internal/runtime/kubernetes/snapshot_progress_test.go +++ /dev/null @@ -1,93 +0,0 @@ -package kubernetes - -import ( - "strings" - "testing" - "time" - - "github.com/highcard-dev/daemon/internal/core/domain" -) - -func TestObserveSnapshotProgressReadsSteamCMDByteMarker(t *testing.T) { - progress := domain.NewSnapshotProgress() - line := `DRUID_PROGRESS_V1 {"step_id":"steamcmd","stage":"downloading","label":"Downloading server files","unit":"bytes","current":12549178426,"total":22938933947,"total_final":true}` - - if !observeSnapshotProgress(line, progress) { - t.Fatal("progress marker was not recognized") - } - if got := progress.Percentage(); got != 54.7 { - t.Fatalf("percentage = %v; want 54.7", got) - } -} - -func TestObserveSnapshotProgressReadsLowProgressByteMarker(t *testing.T) { - progress := domain.NewSnapshotProgress() - line := `DRUID_PROGRESS_V1 {"unit":"bytes","current":401991074,"total":22938933947}` - - if !observeSnapshotProgress(line, progress) { - t.Fatal("progress marker was not recognized") - } - if got := progress.Percentage(); got != 1.8 { - t.Fatalf("percentage = %v; want 1.8", got) - } -} - -func TestObserveSnapshotProgressReadsOriginalSteamCMDLine(t *testing.T) { - progress := domain.NewSnapshotProgress() - line := "\x1b[0m Update state (0x61) downloading, progress: 54.71 (12549178426 / 22938933947)" - - if observeSnapshotProgress(line, progress) { - t.Fatal("original SteamCMD output should remain visible in the console") - } - if got := progress.Percentage(); got != 54.7 { - t.Fatalf("percentage = %v; want 54.7", got) - } -} - -func TestReadSnapshotProgressDoesNotDependOnConsoleConsumer(t *testing.T) { - output := make(chan string, 1) - output <- "older console line" - progress := domain.NewSnapshotProgress() - done := make(chan struct{}) - - go func() { - _ = readSnapshotProgress( - strings.NewReader( - "ordinary console line\n"+ - `DRUID_PROGRESS_V1 {"unit":"bytes","current":12549178426,"total":22938933947}`+"\n", - ), - progress, - ) - close(done) - }() - - select { - case <-done: - case <-time.After(100 * time.Millisecond): - t.Fatal("progress processing blocked behind the console backlog") - } - if got := progress.Percentage(); got != 54.7 { - t.Fatalf("percentage = %v; want 54.7", got) - } - if got := <-output; got != "older console line" { - t.Fatalf("console backlog changed to %q", got) - } -} - -func TestObserveSnapshotProgressLeavesMalformedLinesAlone(t *testing.T) { - progress := domain.NewSnapshotProgress() - progress.StorePercentage(17) - - for _, line := range []string{ - "ordinary server output", - `DRUID_PROGRESS_V1 {"unit":"bytes","current":10,"total":0}`, - `DRUID_PROGRESS_V1 {"unit":"items","current":10,"total":20}`, - } { - if observeSnapshotProgress(line, progress) { - t.Fatalf("line was unexpectedly recognized: %s", line) - } - if got := progress.Percentage(); got != 17 { - t.Fatalf("percentage changed to %v for line %q", got, line) - } - } -} diff --git a/internal/runtime/kubernetes/wait_jobs.go b/internal/runtime/kubernetes/wait_jobs.go index 71436e53..5e1422d1 100644 --- a/internal/runtime/kubernetes/wait_jobs.go +++ b/internal/runtime/kubernetes/wait_jobs.go @@ -14,7 +14,6 @@ import ( metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/apimachinery/pkg/labels" - "github.com/highcard-dev/daemon/internal/core/domain" "github.com/highcard-dev/daemon/internal/utils/logger" "go.uber.org/zap" ) @@ -254,86 +253,48 @@ func (b *Backend) podLogs(ctx context.Context, namespace string, podName string) return logs, nil } -func (b *Backend) streamPodLogs(ctx context.Context, namespace string, podName string, output chan<- string, progress *domain.SnapshotProgress) { +func (b *Backend) streamPodLogs(ctx context.Context, namespace string, podName string, output chan<- string) { defer close(output) - stream, err := b.openPodLogStream(ctx, namespace, podName) - if err != nil { - output <- fmt.Sprintf("failed to stream pod logs: %v", err) - return - } - defer stream.Close() - var progressDone chan struct{} - if progress != nil { - progressDone = make(chan struct{}) - go func() { - defer close(progressDone) - b.streamPodProgress(ctx, namespace, podName, progress) - }() - } - scanner := bufio.NewScanner(stream) - for scanner.Scan() { - line := scanner.Text() - if progress != nil && strings.HasPrefix(line, snapshotProgressPrefix) { - continue - } - output <- line - } - if progressDone != nil { - <-progressDone - } - if err := scanner.Err(); err != nil { - logger.Log().Warn("Kubernetes pod log stream ended with scanner error", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - return - } - logger.Log().Debug("Kubernetes pod log stream ended", zap.String("namespace", namespace), zap.String("pod", podName)) -} - -func (b *Backend) openPodLogStream(ctx context.Context, namespace string, podName string) (io.ReadCloser, error) { + var stream io.ReadCloser deadline := time.Now().Add(30 * time.Second) logger.Log().Debug("Opening Kubernetes follow log stream", zap.String("namespace", namespace), zap.String("pod", podName)) for { req := b.client.CoreV1().Pods(namespace).GetLogs(podName, &corev1.PodLogOptions{Follow: true}) - stream, err := req.Stream(ctx) + var err error + stream, err = req.Stream(ctx) if err == nil { logger.Log().Debug("Kubernetes follow log stream opened", zap.String("namespace", namespace), zap.String("pod", podName)) - return stream, nil + break } if !strings.Contains(err.Error(), "ContainerCreating") && !strings.Contains(err.Error(), "PodInitializing") && !strings.Contains(err.Error(), "not available") { logger.Log().Warn("Failed to stream Kubernetes pod logs", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - return nil, err + output <- fmt.Sprintf("failed to stream pod logs: %v", err) + return } if time.Now().After(deadline) { logger.Log().Warn("Timed out opening Kubernetes pod log stream", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - return nil, err + output <- fmt.Sprintf("failed to stream pod logs: %v", err) + return } logger.Log().Debug("Kubernetes pod logs not ready yet", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) select { case <-ctx.Done(): logger.Log().Warn("Context cancelled while opening Kubernetes pod logs", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(ctx.Err())) - return nil, ctx.Err() + output <- fmt.Sprintf("failed to stream pod logs: %v", ctx.Err()) + return case <-time.After(500 * time.Millisecond): } } -} - -func (b *Backend) streamPodProgress(ctx context.Context, namespace string, podName string, progress *domain.SnapshotProgress) { - stream, err := b.openPodLogStream(ctx, namespace, podName) - if err != nil { - logger.Log().Warn("Failed to stream Kubernetes pod progress", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - return - } defer stream.Close() - if err := readSnapshotProgress(stream, progress); err != nil { - logger.Log().Warn("Kubernetes pod progress stream ended with scanner error", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) - } -} - -func readSnapshotProgress(input io.Reader, progress *domain.SnapshotProgress) error { - scanner := bufio.NewScanner(input) + scanner := bufio.NewScanner(stream) for scanner.Scan() { - observeSnapshotProgress(scanner.Text(), progress) + output <- scanner.Text() + } + if err := scanner.Err(); err != nil { + logger.Log().Warn("Kubernetes pod log stream ended with scanner error", zap.String("namespace", namespace), zap.String("pod", podName), zap.Error(err)) + return } - return scanner.Err() + logger.Log().Debug("Kubernetes pod log stream ended", zap.String("namespace", namespace), zap.String("pod", podName)) }