diff --git a/.agents/distributed-seams.md b/.agents/distributed-seams.md index a14649d8d..2b313ed1f 100644 --- a/.agents/distributed-seams.md +++ b/.agents/distributed-seams.md @@ -48,6 +48,15 @@ moves files. `S3NATSFileStager` returns `nodes.ErrNoRoute` when nothing is listening for the node. `HTTPFileStager` reports connection failures as ordinary errors. +`NodeCommandSender` embeds `LoadOperationControl`: the load operation verbs +(`InstallBackendOp`, `StopLoadOperation`, `OperationControl` for renewals and +completions, `UnloadReplica`, `StopModelReplica`). A carrier implements all of +them. Its timing and error contract is the doc comment on the interface in +`core/services/nodes/interfaces.go`, and `load_operation_control_conformance_test.go` +runs every carrier in `loadOperationCarriers` against it. Whether a worker +names operations is the worker's own report (`BackendInstallReply.ReportsOperations`), +not a property of the carrier. + Worker half: each verb is a `controlVerb`. A handler is typed with `unary`, `withProgress` or `noReply` and registered with `handle` (one request of the verb at a time on NATS, panic not recovered) or `handleWithProgress` (a diff --git a/core/application/application.go b/core/application/application.go index ce333a5a0..326885b31 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -608,6 +608,9 @@ func (a *Application) start() error { assistantClient.StatsRecorder = a.statsRecorder assistantClient.FallbackUser = a.fallbackUser assistantClient.VoiceProfiles = a.voiceProfileStore + if a.distributed != nil && a.distributed.Unloader != nil { + assistantClient.LoadStopper = a.distributed.Unloader + } // PII filter — same nil-or-real wiring. assistantClient.PIIRedactor = a.piiRedactor assistantClient.PIIEvents = a.piiEvents diff --git a/core/http/endpoints/localai/api_instructions.go b/core/http/endpoints/localai/api_instructions.go index 49f2bd0b9..06295b0ef 100644 --- a/core/http/endpoints/localai/api_instructions.go +++ b/core/http/endpoints/localai/api_instructions.go @@ -55,6 +55,7 @@ var instructionDefs = []instructionDef{ }, { Name: "model-management", + Intro: "GET /api/models/{id}/load-status reports job_id and the lease of a distributed load. Admin POST /api/models/{id}/load-cancel needs that job_id: 200 means stopped or gone, 202 means the cancel is recorded and the stop is pending (retry_after says when the model is released regardless), 409 means a different attempt is current. Never cancel a replacement attempt.", Description: "Browse the gallery, install, delete, and manage models and backends", Tags: []string{"models", "backends"}, }, diff --git a/core/http/endpoints/localai/model_load_cancel.go b/core/http/endpoints/localai/model_load_cancel.go new file mode 100644 index 000000000..4d77edf99 --- /dev/null +++ b/core/http/endpoints/localai/model_load_cancel.go @@ -0,0 +1,81 @@ +package localai + +import ( + "encoding/json" + "errors" + "io" + "math" + "net/http" + "strconv" + "strings" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/nodes" +) + +// ModelLoadCancelEndpoint cancels one distributed load attempt. +// +// The body names the exact attempt (job_id, from load-status). That precondition +// is what keeps a cancel from hitting a replacement attempt that started after +// the caller looked. The call is idempotent, and repeating it never extends the +// time the model is held. +// +// @Summary Cancel one distributed model load. +// @Description Cancels the load attempt named by `job_id` and stops its remote work. 200 means the attempt is gone or the worker confirmed the stop. 202 means the cancel is recorded and the stop is pending; the model is released after `retry_after` seconds regardless. 409 means a different attempt is current and carries its `current_job_id`. Repeating the call is safe and does not extend the hold. +// @Tags models +// @Accept json +// @Produce json +// @Param id path string true "Model ID" +// @Param request body schema.ModelLoadCancelRequest true "The exact attempt to cancel" +// @Success 200 {object} schema.ModelLoadCancelResponse "Stopped, or no such load any more" +// @Success 202 {object} schema.ModelLoadCancelResponse "Cancel recorded; stop pending" +// @Failure 400 {object} schema.ErrorResponse +// @Failure 401 {object} schema.ErrorResponse +// @Failure 403 {object} schema.ErrorResponse +// @Failure 404 {object} schema.ErrorResponse "Unknown model" +// @Failure 409 {object} schema.ModelLoadCancelResponse "A different attempt is current" +// @Router /api/models/{id}/load-cancel [post] +func ModelLoadCancelEndpoint(service func() *nodes.LoadCancelService, modelKnown func(id string) bool) echo.HandlerFunc { + return func(c echo.Context) error { + var req schema.ModelLoadCancelRequest + dec := json.NewDecoder(http.MaxBytesReader(c.Response(), c.Request().Body, 4096)) + dec.DisallowUnknownFields() + if err := dec.Decode(&req); err != nil { + return c.JSON(http.StatusBadRequest, nodeError(http.StatusBadRequest, "job_id required in a JSON object")) + } + if strings.TrimSpace(req.JobID) == "" || dec.Decode(new(any)) != io.EOF { + return c.JSON(http.StatusBadRequest, nodeError(http.StatusBadRequest, "invalid cancellation body")) + } + + var svc *nodes.LoadCancelService + if service != nil { + svc = service() + } + model := c.Param("id") + if svc == nil { + // Not distributed: there are no load jobs to cancel. + return c.JSON(http.StatusNotFound, nodeError(http.StatusNotFound, "no distributed loads on this server")) + } + + result, err := svc.Cancel(c.Request().Context(), nodes.LoadJobRef{TrackingKey: model, Generation: req.JobID}) + if errors.Is(err, nodes.ErrLoadCancelConflict) { + return c.JSON(http.StatusConflict, schema.ModelLoadCancelResponse{Model: model, JobID: req.JobID, State: "conflict", CurrentJobID: result.CurrentJobID}) + } + if err != nil { + return c.JSON(http.StatusInternalServerError, nodeError(http.StatusInternalServerError, err.Error())) + } + if result.State == nodes.LoadCancelGone && modelKnown != nil && !modelKnown(model) { + return c.JSON(http.StatusNotFound, nodeError(http.StatusNotFound, "unknown model "+model)) + } + + body := schema.ModelLoadCancelResponse{Model: model, JobID: req.JobID, State: string(result.State)} + code := http.StatusOK + if result.State == nodes.LoadCancelStopping { + code = http.StatusAccepted + body.RetryAfter = int(math.Ceil(result.RetryAfter.Seconds())) + c.Response().Header().Set("Retry-After", strconv.Itoa(body.RetryAfter)) + } + return c.JSON(code, body) + } +} diff --git a/core/http/endpoints/localai/model_load_cancel_test.go b/core/http/endpoints/localai/model_load_cancel_test.go new file mode 100644 index 000000000..2bdc836c2 --- /dev/null +++ b/core/http/endpoints/localai/model_load_cancel_test.go @@ -0,0 +1,313 @@ +package localai_test + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "sync" + + "github.com/labstack/echo/v4" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/http/endpoints/localai" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// opSender is the worker side as the HTTP layer sees it: it can stop one load +// operation, and it records every other call so a spec can tell what was and +// was not sent. +type opSender struct { + nodes.NodeCommandSender // any call a spec does not expect panics + + mu sync.Mutex + hang bool + stops []workerctl.ModelStopRequest + unloads []string + stopBackends []string +} + +func (s *opSender) StopLoadOperation(_ context.Context, _ string, req workerctl.ModelStopRequest) (workerctl.ModelStopReply, error) { + s.mu.Lock() + defer s.mu.Unlock() + s.stops = append(s.stops, req) + if s.hang { + return workerctl.ModelStopReply{}, errors.New("nats: timeout") + } + return workerctl.ModelStopReply{Matched: true, Terminated: true}, nil +} + +func (s *opSender) UnloadModelOnNode(_, model string) error { + s.mu.Lock() + defer s.mu.Unlock() + s.unloads = append(s.unloads, model) + return nil +} + +func (s *opSender) StopBackend(_, model string) error { + s.mu.Lock() + defer s.mu.Unlock() + s.stopBackends = append(s.stopBackends, model) + return nil +} + +var _ = Describe("ModelLoadCancelEndpoint", func() { + var ( + registry *nodes.NodeRegistry + sender *opSender + ctx context.Context + node *nodes.BackendNode + ) + + BeforeEach(func() { + db := testutil.SetupTestDB() + var err error + registry, err = nodes.NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + ctx = context.Background() + sender = &opSender{} + node = &nodes.BackendNode{Name: "worker-1", NodeType: nodes.NodeTypeBackend, Address: "10.0.0.1:50051"} + Expect(registry.Register(ctx, node, true)).To(Succeed()) + }) + + service := func() *nodes.LoadCancelService { + return &nodes.LoadCancelService{Registry: registry, Stopper: sender} + } + serve := func(known func(string) bool) *echo.Echo { + e := echo.New() + e.POST("/api/models/:id/load-cancel", localai.ModelLoadCancelEndpoint(service, known)) + return e + } + post := func(e *echo.Echo, model, body string) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, "/api/models/"+model+"/load-cancel", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + return rec + } + cancelBody := func(jobID string) string { return `{"job_id":"` + jobID + `"}` } + placed := func(model string) *nodes.ModelLoadJob { + job, claimed, err := registry.ClaimLoadJob(ctx, model, "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(registry.UpdateLoadJob(ctx, job.Ref(), nodes.LoadJobUpdate{State: nodes.LoadJobStateLoading, NodeID: node.ID, NodeName: node.Name})).To(Succeed()) + return job + } + known := func(string) bool { return true } + + It("rejects a missing job id, unknown fields and trailing JSON", func() { + e := serve(known) + for _, body := range []string{`{}`, `{"job_id":""}`, `{"job_id":"x","other":true}`, `{"job_id":"x"} {}`, `null`, `nope`} { + Expect(post(e, "m", body).Code).To(Equal(http.StatusBadRequest), body) + } + }) + + It("answers the whole matrix: stopped, stopping, gone, unknown, conflict", func() { + e := serve(func(id string) bool { return id != "unknown" }) + job := placed("m") + + // A worker that confirms: 200 stopped, and the model is released shortly. + rec := post(e, "m", cancelBody(job.Generation)) + Expect(rec.Code).To(Equal(http.StatusOK)) + var body schema.ModelLoadCancelResponse + Expect(json.Unmarshal(rec.Body.Bytes(), &body)).To(Succeed()) + Expect(body.State).To(Equal("stopped")) + Expect(sender.stops).To(HaveLen(1)) + Expect(sender.stops[0].OperationID).To(Equal(job.Generation)) + current, err := registry.GetLoadJob(ctx, "m") + Expect(err).ToNot(HaveOccurred()) + Expect(current.CancelRequested).To(BeTrue()) + Expect(current.State).To(Equal(nodes.LoadJobStateFailed)) + + // A different generation is current: 409 with its id. + rec = post(e, "m", cancelBody("not-the-current-one")) + Expect(rec.Code).To(Equal(http.StatusConflict)) + Expect(json.Unmarshal(rec.Body.Bytes(), &body)).To(Succeed()) + Expect(body.CurrentJobID).To(Equal(job.Generation)) + + // No load at all: 200 gone for a known model, 404 for an unknown one. + rec = post(e, "idle", cancelBody("x")) + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(json.Unmarshal(rec.Body.Bytes(), &body)).To(Succeed()) + Expect(body.State).To(Equal("gone")) + Expect(post(e, "unknown", cancelBody("x")).Code).To(Equal(http.StatusNotFound)) + + // A server that is not distributed has no loads to cancel. + plain := echo.New() + plain.POST("/api/models/:id/load-cancel", localai.ModelLoadCancelEndpoint(func() *nodes.LoadCancelService { return nil }, known)) + Expect(post(plain, "m", cancelBody("x")).Code).To(Equal(http.StatusNotFound)) + }) + + It("answers 202 while the worker is silent, and a repeat does not extend the hold", func() { + sender.hang = true + e := serve(known) + job := placed("silent") + + rec := post(e, "silent", cancelBody(job.Generation)) + Expect(rec.Code).To(Equal(http.StatusAccepted)) + Expect(rec.Header().Get("Retry-After")).ToNot(BeEmpty()) + var body schema.ModelLoadCancelResponse + Expect(json.Unmarshal(rec.Body.Bytes(), &body)).To(Succeed()) + Expect(body.State).To(Equal("stopping")) + Expect(body.RetryAfter).To(BeNumerically(">", 100)) + + first, err := registry.GetLoadJob(ctx, "silent") + Expect(err).ToNot(HaveOccurred()) + Expect(first.StopDeadline).ToNot(BeNil()) + + // Repeating retries the stop and keeps the same deadline. + Expect(post(e, "silent", cancelBody(job.Generation)).Code).To(Equal(http.StatusAccepted)) + second, err := registry.GetLoadJob(ctx, "silent") + Expect(err).ToNot(HaveOccurred()) + Expect(second.StopDeadline.Equal(*first.StopDeadline)).To(BeTrue(), "a repeated cancel must not extend the hold") + Expect(sender.stops).To(HaveLen(2)) + + // The worker answers on the next repeat: the cancel completes. + sender.hang = false + Expect(post(e, "silent", cancelBody(job.Generation)).Code).To(Equal(http.StatusOK)) + }) + + It("cancels a staging load that has no node yet, without sending a stop", func() { + job, _, err := registry.ClaimLoadJob(ctx, "staging", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + + rec := post(serve(known), "staging", cancelBody(job.Generation)) + + Expect(rec.Code).To(Equal(http.StatusAccepted)) + Expect(sender.stops).To(BeEmpty()) + current, err := registry.GetLoadJob(ctx, "staging") + Expect(err).ToNot(HaveOccurred()) + Expect(current.CancelRequested).To(BeTrue()) + }) + + It("is admin only, and not in the quota route registry", func() { + for _, tc := range []struct { + role string + code int + }{{"", http.StatusUnauthorized}, {auth.RoleUser, http.StatusForbidden}, {auth.RoleAdmin, http.StatusBadRequest}} { + e := echo.New() + e.Use(func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + if tc.role != "" { + c.Set("auth_user", &auth.User{Role: tc.role}) + } + return next(c) + } + }) + e.POST("/api/models/:id/load-cancel", localai.ModelLoadCancelEndpoint(service, known), auth.RequireAdmin()) + Expect(post(e, "m", `{}`).Code).To(Equal(tc.code), tc.role) + } + for _, entry := range auth.RouteFeatureRegistry { + Expect(entry.Pattern).ToNot(ContainSubstring("load-cancel"), "a management route must not be metered as a modality") + } + }) + + Describe("unloading and deregistering a node", func() { + It("cancels the in-flight load and still unloads the loaded replica", func() { + job := placed("both") + Expect(registry.SetNodeModel(ctx, node.ID, "both", 1, "loaded", "10.0.0.1:9001", 0)).To(Succeed()) + other := placed("elsewhere") + Expect(registry.UpdateLoadJob(ctx, other.Ref(), nodes.LoadJobUpdate{NodeID: "another-node", NodeName: "n2"})).To(Succeed()) + + e := echo.New() + e.POST("/api/nodes/:id/models/unload", localai.UnloadModelOnNodeEndpoint(sender, registry)) + req := httptest.NewRequest(http.MethodPost, "/api/nodes/"+node.ID+"/models/unload", strings.NewReader(`{"model_name":"both"}`)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(sender.stops).To(HaveLen(1)) + Expect(sender.stops[0].OperationID).To(Equal(job.Generation)) + Expect(sender.unloads).To(ConsistOf("both"), "the loaded replica is still unloaded") + Expect(sender.stopBackends).To(ConsistOf("both")) + replicas, err := registry.GetNodeModels(ctx, node.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(replicas).To(BeEmpty()) + cancelled, err := registry.GetLoadJob(ctx, "both") + Expect(err).ToNot(HaveOccurred()) + Expect(cancelled.CancelRequested).To(BeTrue()) + + // A load of another model, placed on another node, is untouched. + untouched, err := registry.GetLoadJob(ctx, "elsewhere") + Expect(err).ToNot(HaveOccurred()) + Expect(untouched.State).To(Equal(nodes.LoadJobStateLoading)) + }) + + It("stops the operations of a node before deregistering it", func() { + job := placed("on-node") + + e := echo.New() + e.DELETE("/api/nodes/:id", localai.DeregisterNodeEndpoint(registry, sender)) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodDelete, "/api/nodes/"+node.ID, nil)) + + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(sender.stops).To(HaveLen(1)) + Expect(sender.stops[0].OperationID).To(Equal(job.Generation)) + removed, err := registry.GetLoadJob(ctx, "on-node") + Expect(err).ToNot(HaveOccurred()) + Expect(removed.CancelRequested).To(BeFalse(), "no administrator cancelled this load") + Expect(removed.LastError).ToNot(ContainSubstring("administrator")) + }) + }) +}) + +var _ = Describe("ModelLoadStatusEndpoint activity", func() { + It("projects the job id, the lease, and for a failed load the cause and the hold", func() { + db := testutil.SetupTestDB() + registry, err := nodes.NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + ctx := context.Background() + e := echo.New() + e.GET("/api/models/:id/load-status", localai.ModelLoadStatusEndpoint(func() nodes.LoadJobStore { return registry })) + get := func(model string) (int, schema.ModelLoadingStatus) { + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/models/"+model+"/load-status", nil)) + var status schema.ModelLoadingStatus + _ = json.Unmarshal(rec.Body.Bytes(), &status) + return rec.Code, status + } + + job, _, err := registry.ClaimLoadJob(ctx, "m", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + code, status := get("m") + Expect(code).To(Equal(http.StatusOK)) + Expect(status.JobID).To(Equal(job.Generation)) + Expect(status.LeaseExpiresIn).ToNot(BeNil()) + Expect(*status.LeaseExpiresIn).To(BeNumerically(">", 0)) + Expect(status.Stopping).To(BeFalse()) + + Expect(registry.FailLoadJob(ctx, job.Ref(), "context deadline exceeded", true)).To(Succeed()) + code, status = get("m") + Expect(code).To(Equal(http.StatusOK)) + Expect(status.State).To(Equal(nodes.LoadJobStateFailed)) + Expect(status.LastError).To(Equal("context deadline exceeded")) + Expect(status.Stopping).To(BeTrue(), "the remote work is not confirmed ended") + Expect(status.StopDeadline).ToNot(BeNil()) + Expect(status.RetryAfter).To(BeNumerically(">", 100)) + Expect(status.LeaseExpiresIn).To(BeNil()) + }) + + It("fails closed with 503 when the job table cannot be read", func() { + e := echo.New() + e.GET("/api/models/:id/load-status", localai.ModelLoadStatusEndpoint(func() nodes.LoadJobStore { return brokenStore{} })) + rec := httptest.NewRecorder() + e.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/models/m/load-status", nil)) + Expect(rec.Code).To(Equal(http.StatusServiceUnavailable)) + }) +}) + +// brokenStore fails every read, like a database that is down. +type brokenStore struct{ nodes.LoadJobStore } + +func (brokenStore) GetLoadJob(context.Context, string) (*nodes.ModelLoadJob, error) { + return nil, errors.New("database unreachable") +} diff --git a/core/http/endpoints/localai/model_load_status.go b/core/http/endpoints/localai/model_load_status.go index 70cebdcaa..ae2d75a16 100644 --- a/core/http/endpoints/localai/model_load_status.go +++ b/core/http/endpoints/localai/model_load_status.go @@ -18,12 +18,13 @@ import ( // for a 503 depend on which modality the model happens to be. // // @Summary Report the progress of an in-flight model load. -// @Description Returns the live state of a distributed cold load — phase, node, byte progress and ETA — or 404 when no load is running for the model. This is the same `loading` object the 503 response carries while a model is still staging. +// @Description Returns the live state of a distributed cold load: job id, phase, node, byte progress, ETA, lease freshness, and for a failed attempt the cause, whether the stop is pending and when the model is released. 404 means no load exists for the model. A database error is 503, never an empty answer. This is the same `loading` object the 503 response carries while a model is still staging. // @Tags models // @Produce json // @Param id path string true "Model ID" // @Success 200 {object} schema.ModelLoadingStatus "Live load progress" // @Failure 404 {object} schema.ErrorResponse "No load is running for this model" +// @Failure 503 {object} schema.ErrorResponse "The job table could not be read" // @Router /api/models/{id}/load-status [get] func ModelLoadStatusEndpoint(loadJobs func() nodes.LoadJobStore) echo.HandlerFunc { return func(c echo.Context) error { @@ -54,8 +55,10 @@ func ModelLoadStatusEndpoint(loadJobs func() nodes.LoadJobStore) echo.HandlerFun job, err := store.GetLoadJob(c.Request().Context(), modelID) if err != nil { - return c.JSON(http.StatusInternalServerError, schema.ErrorResponse{ - Error: &schema.APIError{Message: err.Error(), Code: http.StatusInternalServerError, Type: "server_error"}, + // Fail closed: a database error must not read as "no load is + // running", or a caller would start a second one. + return c.JSON(http.StatusServiceUnavailable, schema.ErrorResponse{ + Error: &schema.APIError{Message: err.Error(), Code: http.StatusServiceUnavailable, Type: "server_error"}, }) } if job == nil { diff --git a/core/http/endpoints/localai/model_load_status_test.go b/core/http/endpoints/localai/model_load_status_test.go index 9f989a548..35c69914d 100644 --- a/core/http/endpoints/localai/model_load_status_test.go +++ b/core/http/endpoints/localai/model_load_status_test.go @@ -52,10 +52,10 @@ var _ = Describe("ModelLoadStatusEndpoint", func() { It("reports the live progress of a running load", func() { ctx := context.Background() - _, claimed, err := registry.ClaimLoadJob(ctx, "big-model", "replica-a") + job, claimed, err := registry.ClaimLoadJob(ctx, "big-model", "replica-a") Expect(err).ToNot(HaveOccurred()) Expect(claimed).To(BeTrue()) - Expect(registry.UpdateLoadJob(ctx, "big-model", nodes.LoadJobUpdate{ + Expect(registry.UpdateLoadJob(ctx, job.Ref(), nodes.LoadJobUpdate{ State: nodes.LoadJobStateStaging, NodeID: "n1", NodeName: "nvidia-thor", BytesSent: 1000, TotalBytes: 4000, FileIndex: 1, TotalFiles: 1, })).To(Succeed()) diff --git a/core/http/endpoints/localai/nodes.go b/core/http/endpoints/localai/nodes.go index 32c0ffe4e..b38cb5a60 100644 --- a/core/http/endpoints/localai/nodes.go +++ b/core/http/endpoints/localai/nodes.go @@ -356,11 +356,33 @@ func provisionAgentWorkerKey(ctx context.Context, authDB *gorm.DB, registry *nod return plaintext, nil } +// loadCancelService builds the cancel path unload, drain and deregister share +// with the load-cancel endpoint. A sender that cannot stop operations still +// gets the cancel recorded; the worker's watchdog and the stop window then +// bound the work. +func loadCancelService(registry *nodes.NodeRegistry, unloader nodes.NodeCommandSender) *nodes.LoadCancelService { + svc := &nodes.LoadCancelService{Registry: registry} + if unloader != nil { + svc.Stopper = unloader + } + return svc +} + +// cancelNodeLoads cancels the loads placed on a node before its rows go. It +// logs and carries on: a node being removed must not stay because a cancel +// failed. +func cancelNodeLoads(ctx context.Context, registry *nodes.NodeRegistry, unloader nodes.NodeCommandSender, nodeID string) { + if err := loadCancelService(registry, unloader).CancelNodeLoads(ctx, nodeID); err != nil { + xlog.Warn("Failed to cancel the loads placed on a node", "node", nodeID, "error", err) + } +} + // DeregisterNodeEndpoint removes a backend node permanently (admin use). -func DeregisterNodeEndpoint(registry *nodes.NodeRegistry) echo.HandlerFunc { +func DeregisterNodeEndpoint(registry *nodes.NodeRegistry, unloader nodes.NodeCommandSender) echo.HandlerFunc { return func(c echo.Context) error { ctx := c.Request().Context() id := c.Param("id") + cancelNodeLoads(ctx, registry, unloader, id) if err := registry.Deregister(ctx, id); err != nil { xlog.Error("Failed to deregister node", "id", id, "error", err) return c.JSON(http.StatusInternalServerError, nodeError(http.StatusInternalServerError, "failed to deregister node")) @@ -400,6 +422,11 @@ func HeartbeatEndpoint(registry *nodes.NodeRegistry) echo.HandlerFunc { } ctx := c.Request().Context() + // A new incarnation means the worker restarted and every load operation of + // the previous process ended. Failures here must not fail the heartbeat. + if err := registry.ObserveWorkerIncarnation(ctx, id, update.WorkerIncarnation); err != nil { + xlog.Warn("Failed to record worker incarnation", "id", id, "error", err) + } if err := registry.Heartbeat(ctx, id, updatePtr); err != nil { xlog.Warn("Heartbeat failed for node", "id", id, "error", err) return c.JSON(http.StatusNotFound, nodeError(http.StatusNotFound, "node not found")) @@ -440,10 +467,11 @@ func ListAllNodeModelsEndpoint(registry *nodes.NodeRegistry) echo.HandlerFunc { } // DrainNodeEndpoint sets a node to draining status (no new requests). -func DrainNodeEndpoint(registry *nodes.NodeRegistry) echo.HandlerFunc { +func DrainNodeEndpoint(registry *nodes.NodeRegistry, unloader nodes.NodeCommandSender) echo.HandlerFunc { return func(c echo.Context) error { ctx := c.Request().Context() id := c.Param("id") + cancelNodeLoads(ctx, registry, unloader, id) if err := registry.MarkDraining(ctx, id); err != nil { if errors.Is(err, nodes.ErrNodeNotFound) { return c.JSON(http.StatusNotFound, nodeError(http.StatusNotFound, "node not found")) @@ -700,6 +728,12 @@ func UnloadModelOnNodeEndpoint(unloader nodes.NodeCommandSender, registry *nodes if err := c.Bind(&req); err != nil || req.ModelName == "" { return c.JSON(http.StatusBadRequest, nodeError(http.StatusBadRequest, "model_name required")) } + // A load of this model on this node (or one not placed yet) is cancelled + // through the same stop path first. The unload of a loaded replica then + // proceeds as usual: it is no longer swallowed by the load. + if _, err := loadCancelService(registry, unloader).CancelModelOnNode(c.Request().Context(), nodeID, req.ModelName); err != nil { + xlog.Warn("Failed to cancel the load before unloading", "node", nodeID, "model", req.ModelName, "error", err) + } if err := unloader.UnloadModelOnNode(nodeID, req.ModelName); err != nil { xlog.Error("Failed to unload model on node", "node", nodeID, "model", req.ModelName, "error", err) return c.JSON(http.StatusInternalServerError, nodeError(http.StatusInternalServerError, "failed to unload model on node")) diff --git a/core/http/endpoints/localai/nodes_backends_list_test.go b/core/http/endpoints/localai/nodes_backends_list_test.go index b9f0a0640..bc9fa918d 100644 --- a/core/http/endpoints/localai/nodes_backends_list_test.go +++ b/core/http/endpoints/localai/nodes_backends_list_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "time" "github.com/labstack/echo/v4" "github.com/mudler/LocalAI/core/services/nodes" @@ -44,6 +45,25 @@ func (s *stubNodeCommandSender) UnloadModelOnNode(_, _ string) error { return ni func (s *stubNodeCommandSender) PingNode(_ string) error { return nil } +// The load operation verbs. This stub never reaches them. +func (s *stubNodeCommandSender) InstallBackendOp(_, _, _, _ string, _ int, _, _ string, _ time.Duration, _ func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { + return &workerctl.BackendInstallReply{}, nil +} + +func (s *stubNodeCommandSender) StopLoadOperation(context.Context, string, workerctl.ModelStopRequest) (workerctl.ModelStopReply, error) { + return workerctl.ModelStopReply{}, nil +} + +func (s *stubNodeCommandSender) OperationControl(string, workerctl.OperationRequest) (*workerctl.OperationReply, error) { + return &workerctl.OperationReply{}, nil +} + +func (s *stubNodeCommandSender) UnloadReplica(string, nodes.NodeModel) error { return nil } + +func (s *stubNodeCommandSender) StopModelReplica(context.Context, string, nodes.NodeModel, bool) (workerctl.ModelStopReply, error) { + return workerctl.ModelStopReply{}, nil +} + var _ = Describe("ListBackendsOnNodeEndpoint", func() { var registry *nodes.NodeRegistry diff --git a/core/http/endpoints/localai/nodes_test.go b/core/http/endpoints/localai/nodes_test.go index 8390f1a48..736f4bde5 100644 --- a/core/http/endpoints/localai/nodes_test.go +++ b/core/http/endpoints/localai/nodes_test.go @@ -495,7 +495,7 @@ var _ = Describe("Node HTTP handlers", func() { ID: "lifecycle", Name: "lifecycle", Address: "10.0.0.10:50051", }, true)).To(Succeed()) - Expect(request(DrainNodeEndpoint(registry), "lifecycle").Code).To(Equal(http.StatusOK)) + Expect(request(DrainNodeEndpoint(registry, nil), "lifecycle").Code).To(Equal(http.StatusOK)) Expect(request(ResumeNodeEndpoint(registry), "lifecycle").Code).To(Equal(http.StatusOK)) }) @@ -504,12 +504,12 @@ var _ = Describe("Node HTTP handlers", func() { ID: "pending-lifecycle", Name: "pending-lifecycle", Address: "10.0.0.11:50051", }, false)).To(Succeed()) - Expect(request(DrainNodeEndpoint(registry), "pending-lifecycle").Code).To(Equal(http.StatusConflict)) + Expect(request(DrainNodeEndpoint(registry, nil), "pending-lifecycle").Code).To(Equal(http.StatusConflict)) Expect(request(ResumeNodeEndpoint(registry), "pending-lifecycle").Code).To(Equal(http.StatusConflict)) }) It("returns not found for missing nodes", func() { - Expect(request(DrainNodeEndpoint(registry), "missing").Code).To(Equal(http.StatusNotFound)) + Expect(request(DrainNodeEndpoint(registry, nil), "missing").Code).To(Equal(http.StatusNotFound)) Expect(request(ResumeNodeEndpoint(registry), "missing").Code).To(Equal(http.StatusNotFound)) }) }) diff --git a/core/http/endpoints/mcp/localai_assistant_test.go b/core/http/endpoints/mcp/localai_assistant_test.go index 8629bd9e1..3e539ce7d 100644 --- a/core/http/endpoints/mcp/localai_assistant_test.go +++ b/core/http/endpoints/mcp/localai_assistant_test.go @@ -54,6 +54,9 @@ func (stubClient) ReloadModels(_ context.Context) error { return nil } func (stubClient) LoadModel(_ context.Context, model string) ([]string, error) { return []string{model}, nil } +func (stubClient) CancelModelLoad(_ context.Context, model, jobID string) (localaitools.LoadCancelResult, error) { + return localaitools.LoadCancelResult{Model: model, JobID: jobID, State: "stopping"}, nil +} func (stubClient) SetAlias(_ context.Context, _, _ string) error { return nil } diff --git a/core/http/model_loading_held_test.go b/core/http/model_loading_held_test.go new file mode 100644 index 000000000..7eda82db1 --- /dev/null +++ b/core/http/model_loading_held_test.go @@ -0,0 +1,40 @@ +package http + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "time" + + "github.com/labstack/echo/v4" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/nodes" +) + +// A request for a model held by a failed load gets 503 with a Retry-After that +// says when the hold ends, and the real cause. It never gets a bare 500. +var _ = Describe("A model held by a failed load", func() { + It("answers 503 with Retry-After, the cause and the job", func() { + stop := time.Now().Add(120 * time.Second) + job := &nodes.ModelLoadJob{TrackingKey: "held-model", Generation: "gen-1", State: nodes.LoadJobStateFailed, + LastError: "worker out of disk", StopDeadline: &stop} + err := fmt.Errorf("routing: %w", nodes.NewLoadHeldError(job)) + + rec := httptest.NewRecorder() + c := echo.New().NewContext(httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil), rec) + + Expect(respondModelLoading(err, c)).To(BeTrue(), "a held model must be answered as a loading error") + Expect(rec.Code).To(Equal(http.StatusServiceUnavailable)) + Expect(rec.Header().Get("Retry-After")).To(MatchRegexp(`^(1[0-9][0-9]|2[0-9][0-9])$`), "the seconds until the hold ends") + var body schema.ModelLoadingResponse + Expect(json.Unmarshal(rec.Body.Bytes(), &body)).To(Succeed()) + Expect(body.Error).ToNot(BeNil()) + Expect(body.Loading).ToNot(BeNil()) + Expect(body.Loading.LastError).To(Equal("worker out of disk")) + Expect(body.Loading.JobID).To(Equal("gen-1")) + }) +}) diff --git a/core/http/routes/localai.go b/core/http/routes/localai.go index b2bfed146..e6f09f8dc 100644 --- a/core/http/routes/localai.go +++ b/core/http/routes/localai.go @@ -171,6 +171,19 @@ func RegisterLocalAIRoutes(router *echo.Echo, return nil })) + // Cancel one load attempt. Admin only. It is deliberately not in the auth + // RouteFeatureRegistry: that registry meters modality use, and a cancel is + // management, not inference. + router.POST("/api/models/:id/load-cancel", localai.ModelLoadCancelEndpoint(func() *nodes.LoadCancelService { + if d := app.Distributed(); d != nil && d.Registry != nil { + return &nodes.LoadCancelService{Registry: d.Registry, Stopper: d.Unloader} + } + return nil + }, func(id string) bool { + _, ok := cl.GetModelConfig(id) + return ok + }), adminMiddleware) + // Failover chains: reads and the event stream use standard auth (any // authenticated caller may watch chain health); pin/unpin are admin-only // since they override the routing decision for every caller of the chain. @@ -344,6 +357,7 @@ func RegisterLocalAIRoutes(router *echo.Echo, "autocomplete": "/api/models/config-metadata/autocomplete/:provider", "vram_estimate": "/api/models/vram-estimate", "model_load_status": "/api/models/:id/load-status", + "model_load_cancel": "/api/models/:id/load-cancel", "tts": "/tts", "tts_voices": "/v1/audio/voices", "voice_profiles": "/api/voice-profiles", @@ -382,6 +396,7 @@ func RegisterLocalAIRoutes(router *echo.Echo, "reload": "/models/reload", "list_aliases": "/api/aliases", "load_status": "/api/models/:id/load-status", + "load_cancel": "/api/models/:id/load-cancel", }, "ai_functions": map[string]string{ "tts": "/tts", diff --git a/core/http/routes/nodes.go b/core/http/routes/nodes.go index eeb2819e9..2e55cd10b 100644 --- a/core/http/routes/nodes.go +++ b/core/http/routes/nodes.go @@ -47,7 +47,7 @@ func RegisterNodeSelfServiceRoutes(e *echo.Echo, registry *nodes.NodeRegistry, r node := e.Group("/api/node", readyMw, tokenAuthMw) node.POST("/register", localai.RegisterNodeEndpoint(registry, registrationToken, autoApprove, authDB, hmacSecret, natsCfg)) node.POST("/:id/heartbeat", localai.HeartbeatEndpoint(registry)) - node.POST("/:id/drain", localai.DrainNodeEndpoint(registry)) + node.POST("/:id/drain", localai.DrainNodeEndpoint(registry, nil)) node.POST("/:id/resume", localai.ResumeNodeEndpoint(registry)) node.POST("/:id/deregister", localai.DeactivateNodeEndpoint(registry)) node.GET("/:id/models", localai.GetNodeModelsEndpoint(registry)) @@ -82,8 +82,8 @@ func RegisterNodeAdminRoutes(e *echo.Echo, registry *nodes.NodeRegistry, unloade admin.GET("/:id", localai.GetNodeEndpoint(registry)) admin.GET("/:id/models", localai.GetNodeModelsEndpoint(registry)) - admin.DELETE("/:id", localai.DeregisterNodeEndpoint(registry)) - admin.POST("/:id/drain", localai.DrainNodeEndpoint(registry)) + admin.DELETE("/:id", localai.DeregisterNodeEndpoint(registry, unloader)) + admin.POST("/:id/drain", localai.DrainNodeEndpoint(registry, unloader)) admin.POST("/:id/resume", localai.ResumeNodeEndpoint(registry)) admin.POST("/:id/approve", localai.ApproveNodeEndpoint(registry, authDB, hmacSecret, natsCfg)) diff --git a/core/http/routes/ui_api_operations_test.go b/core/http/routes/ui_api_operations_test.go index 28ffc0d73..8b11092b9 100644 --- a/core/http/routes/ui_api_operations_test.go +++ b/core/http/routes/ui_api_operations_test.go @@ -93,15 +93,15 @@ var _ = Describe("/api/operations with durable staging jobs", func() { Expect(err).ToNot(HaveOccurred()) router := nodes.NewSmartRouter(registry, nodes.SmartRouterOptions{}) - _, _, err = registry.ClaimLoadJob(context.Background(), "durable-model", "replica-a") + durableJob, _, err := registry.ClaimLoadJob(context.Background(), "durable-model", "replica-a") Expect(err).ToNot(HaveOccurred()) - Expect(registry.UpdateLoadJob(context.Background(), "durable-model", nodes.LoadJobUpdate{ + Expect(registry.UpdateLoadJob(context.Background(), durableJob.Ref(), nodes.LoadJobUpdate{ State: nodes.LoadJobStateStaging, NodeID: "node-1", NodeName: "durable-node", BytesSent: 25, TotalBytes: 100, FileIndex: 1, TotalFiles: 1, })).To(Succeed()) - _, _, err = registry.ClaimLoadJob(context.Background(), "loading-model", "replica-a") + loadingJob, _, err := registry.ClaimLoadJob(context.Background(), "loading-model", "replica-a") Expect(err).ToNot(HaveOccurred()) - Expect(registry.UpdateLoadJob(context.Background(), "loading-model", nodes.LoadJobUpdate{ + Expect(registry.UpdateLoadJob(context.Background(), loadingJob.Ref(), nodes.LoadJobUpdate{ State: nodes.LoadJobStateLoading, })).To(Succeed()) @@ -125,9 +125,9 @@ var _ = Describe("/api/operations with durable staging jobs", func() { registry, err := nodes.NewNodeRegistry(db) Expect(err).ToNot(HaveOccurred()) router := nodes.NewSmartRouter(registry, nodes.SmartRouterOptions{}) - _, _, err = registry.ClaimLoadJob(context.Background(), "overlay-model", "replica-a") + overlayJob, _, err := registry.ClaimLoadJob(context.Background(), "overlay-model", "replica-a") Expect(err).ToNot(HaveOccurred()) - Expect(registry.UpdateLoadJob(context.Background(), "overlay-model", nodes.LoadJobUpdate{ + Expect(registry.UpdateLoadJob(context.Background(), overlayJob.Ref(), nodes.LoadJobUpdate{ State: nodes.LoadJobStateStaging, NodeName: "durable-node", BytesSent: 10, TotalBytes: 100, })).To(Succeed()) router.StagingTracker().Start("overlay-model", "fresh-node", 1) diff --git a/core/schema/model_loading.go b/core/schema/model_loading.go index 3e2927508..542716444 100644 --- a/core/schema/model_loading.go +++ b/core/schema/model_loading.go @@ -1,10 +1,28 @@ package schema +import "time" + // ModelLoadingStatus describes a cold load that is still in progress. In // distributed mode a model can take tens of minutes to stage onto a worker, // which is far longer than a request may be held; a caller that runs out of // wait budget gets this instead of an anonymous hang or a misleading error. type ModelLoadingStatus struct { + // JobID names the load attempt. A cancel must quote it: it is the + // precondition that keeps a cancel from hitting a replacement attempt. + JobID string `json:"job_id,omitempty"` + // LeaseExpiresIn is the seconds left on the owner's lease. It is negative + // when the lease already ran out, which means the owner is gone. + LeaseExpiresIn *int `json:"lease_expires_in,omitempty"` + // CancelRequested is true when an administrator cancelled the attempt. + CancelRequested bool `json:"cancel_requested,omitempty"` + // LastError is the cause of a failed attempt. + LastError string `json:"last_error,omitempty"` + // Stopping is true while the remote work of a failed attempt is not yet + // confirmed ended. StopDeadline is when the model is released regardless. + Stopping bool `json:"stopping,omitempty"` + StopDeadline *time.Time `json:"stop_deadline,omitempty"` + // RetryAfter is the seconds until a new load may start, for a failed attempt. + RetryAfter int `json:"retry_after,omitempty"` Model string `json:"model"` State string `json:"state"` Node string `json:"node,omitempty"` @@ -26,3 +44,21 @@ type ModelLoadingResponse struct { Error *APIError `json:"error,omitempty"` Loading *ModelLoadingStatus `json:"loading,omitempty"` } + +// ModelLoadCancelRequest is the body of POST /api/models/{id}/load-cancel. +// JobID is the exact attempt to cancel, as load-status reports it. +type ModelLoadCancelRequest struct { + JobID string `json:"job_id"` +} + +// ModelLoadCancelResponse reports what a cancel did. State is "stopped" (the +// worker confirmed the work ended), "stopping" (the cancel is recorded and the +// stop is pending; RetryAfter says when the model is released regardless) or +// "gone" (no such load exists any more). CurrentJobID is set on a conflict. +type ModelLoadCancelResponse struct { + Model string `json:"model"` + JobID string `json:"job_id"` + State string `json:"state"` + RetryAfter int `json:"retry_after,omitempty"` + CurrentJobID string `json:"current_job_id,omitempty"` +} diff --git a/core/services/messaging/subject_rules_test.go b/core/services/messaging/subject_rules_test.go index 09549b152..9f9ac0b64 100644 --- a/core/services/messaging/subject_rules_test.go +++ b/core/services/messaging/subject_rules_test.go @@ -85,7 +85,7 @@ var _ = Describe("Subject rules", func() { messaging.SubjectCacheInvalidateCollection("c1"), messaging.SubjectSyncStateDelta("s1"), messaging.SubjectNodeBackendInstall(id), messaging.SubjectNodeBackendUpgrade(id), messaging.SubjectNodeBackendList(id), messaging.SubjectNodeBackendStop(id), - messaging.SubjectNodeModelStop(id), messaging.SubjectNodeBackendDelete(id), + messaging.SubjectNodeModelStop(id), messaging.SubjectNodeModelOp(id), messaging.SubjectNodeBackendDelete(id), messaging.SubjectNodeModelUnload(id), messaging.SubjectNodeModelDelete(id), messaging.SubjectNodeModelsRunning(id), messaging.SubjectNodeStop(id), messaging.SubjectNodeFilesEnsure(id), messaging.SubjectNodeFilesStage(id), diff --git a/core/services/messaging/subjects.go b/core/services/messaging/subjects.go index 44e475229..7205dd897 100644 --- a/core/services/messaging/subjects.go +++ b/core/services/messaging/subjects.go @@ -194,6 +194,14 @@ func SubjectNodeModelStop(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".model.stop" } +// SubjectNodeModelOp renews and completes the load operations a worker is +// watching. Request-reply, answered with a workerctl.OperationReply. A worker +// that predates it never answers, and the controller then treats the node as +// legacy. +func SubjectNodeModelOp(nodeID string) string { + return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".model.op" +} + // SubjectNodeBackendDelete tells a worker node to delete a backend (stop + remove files). // Uses NATS request-reply. func SubjectNodeBackendDelete(nodeID string) string { diff --git a/core/services/nodes/interfaces.go b/core/services/nodes/interfaces.go index 41e175e0b..52ca43215 100644 --- a/core/services/nodes/interfaces.go +++ b/core/services/nodes/interfaces.go @@ -72,9 +72,89 @@ type ModelRouter interface { type LoadJobStore interface { ClaimLoadJob(ctx context.Context, trackingKey, owner string) (*ModelLoadJob, bool, error) GetLoadJob(ctx context.Context, trackingKey string) (*ModelLoadJob, error) - UpdateLoadJob(ctx context.Context, trackingKey string, u LoadJobUpdate) error - FailLoadJob(ctx context.Context, trackingKey, msg string) error - DeleteLoadJob(ctx context.Context, trackingKey string) error + UpdateLoadJob(ctx context.Context, ref LoadJobRef, u LoadJobUpdate) error + FailLoadJob(ctx context.Context, ref LoadJobRef, msg string, workMayRun bool) error + DeleteLoadJob(ctx context.Context, ref LoadJobRef) error + DeleteFailedLoadJob(ctx context.Context, ref LoadJobRef) error + ConfirmLoadOp(ctx context.Context, ref LoadJobRef) error +} + +// ReplicaUnloader unloads one replica by the address its row recorded. The +// eviction and scale-down paths use it: they delete the row first, so the +// address must travel with the call. +type ReplicaUnloader interface { + UnloadReplica(nodeID string, replica NodeModel) error +} + +// LoadOperationInstaller starts a backend as a load operation the worker bounds. +type LoadOperationInstaller interface { + InstallBackendOp(nodeID, backendType, modelID, galleriesJSON string, replicaIndex int, opID, operationID string, deadline time.Duration, onProgress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) +} + +// LoadOperationStopper is the only path that kills remote load work. It stops +// one operation by id and never "any running backend". +type LoadOperationStopper interface { + StopLoadOperation(ctx context.Context, nodeID string, req workerctl.ModelStopRequest) (workerctl.ModelStopReply, error) +} + +// LoadOperationRenewer renews and completes load operations on a worker node. +type LoadOperationRenewer interface { + OperationControl(nodeID string, req workerctl.OperationRequest) (*workerctl.OperationReply, error) +} + +// LoadAttemptStopper is what stopping a failed or cancelled attempt needs: the +// operation stop for a worker that names operations, and the exact-address stop +// for a worker that does not. +type LoadAttemptStopper interface { + LoadOperationStopper + ExactModelStopper +} + +// LoadOperationControl is the part of a carrier that bounds, renews and stops +// remote load work. Every carrier of NodeCommandSender must provide it; there +// is no degraded mode in which a sender silently lacks it. A worker that cannot +// name operations is a fact about that worker, reported by its install reply +// (BackendInstallReply.ReportsOperations), and is handled per node. +// +// What a carrier owes the callers: +// +// - Every method is a request and a reply with its own timeout. A call that +// gets no reply in time returns an error. It must never block for longer +// than its documented bound, because the owner loop and the reconciler call +// these from timers. +// +// - OperationControl carries renewals and completions. Callers send a renewal +// every 5 seconds for each running load. The worker kills an operation that +// goes 90 seconds without one, so a carrier must deliver a renewal within +// one cadence (5 s, bounded by the 5 s call timeout) or return an error +// promptly. A lost renewal costs nothing until the kill TTL. A completion is +// retried by the caller and its loss is not fatal. +// +// - StopLoadOperation is idempotent. It addresses one operation by operation +// id, process key, process instance and address, and the worker refuses +// unless the ones given match its own records. A repeat after a stop +// returns a reply with Terminated set, not an error. A reply with Error set +// is the worker's refusal: the worker is present and said no. +// +// - InstallBackendOp is InstallBackend that names the operation and its +// longest run as a duration. The deadline is relative so clocks do not +// matter. +// +// - UnloadReplica names the replica's address. With no address nothing is +// sent: a carrier must never ask a worker to pick a process. +// +// - Errors keep the four conditions apart (see .agents/distributed-seams.md). +// ErrNoRoute is a routing fact only: no route from here right now. A +// timeout is not ErrNoRoute. An unreachable peer is not a verdict about the +// worker. A worker's own answer, including a refusal, is a nil error with +// the refusal in the reply. Only the carrier maps its own sentinel onto +// ErrNoRoute. +type LoadOperationControl interface { + LoadOperationInstaller + LoadOperationStopper + LoadOperationRenewer + ReplicaUnloader + ExactModelStopper } // ConcurrencyConflictResolver returns the names of configured models that diff --git a/core/services/nodes/load_cancel.go b/core/services/nodes/load_cancel.go new file mode 100644 index 000000000..11dd2d376 --- /dev/null +++ b/core/services/nodes/load_cancel.go @@ -0,0 +1,223 @@ +package nodes + +import ( + "context" + "errors" + "fmt" + "time" + + "github.com/mudler/LocalAI/core/services/workerctl" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/xlog" +) + +// LoadCancelState is what a cancel achieved. +type LoadCancelState string + +const ( + // LoadCancelStopped: the worker confirmed the load's work ended. + LoadCancelStopped LoadCancelState = "stopped" + // LoadCancelStopping: the cancel is recorded and the stop is pending. The + // model is released at the stop deadline regardless. + LoadCancelStopping LoadCancelState = "stopping" + // LoadCancelGone: no load exists for the model any more. + LoadCancelGone LoadCancelState = "gone" +) + +// ErrLoadCancelConflict means another generation holds the model. +// CurrentJobID on the result names it. +var ErrLoadCancelConflict = errors.New("a different load generation is current") + +// LoadCancelResult is the outcome of one cancel. +type LoadCancelResult struct { + State LoadCancelState + // RetryAfter is how long until the model is released if the stop never + // confirms. Zero when the state is stopped or gone. + RetryAfter time.Duration + // CurrentJobID is the generation that holds the model, set with + // ErrLoadCancelConflict. + CurrentJobID string +} + +// LoadCancelService cancels distributed loads. It is the one path that cancel, +// unload and node deregistration share: it records the cancel on the job, then +// stops the remote work through the single stop call. +type LoadCancelService struct { + Registry *NodeRegistry + // Stopper stops one load operation. It may be nil: the cancel is then + // recorded and the worker's own watchdog and the stop window bound the work. + Stopper LoadAttemptStopper +} + +// Cancel cancels the load attempt ref names. It is idempotent: repeating it on +// a failed attempt retries the stop and never extends the stop window. +func (s *LoadCancelService) Cancel(ctx context.Context, ref LoadJobRef) (LoadCancelResult, error) { + return s.cancel(ctx, ref, "cancelled by an administrator", true) +} + +func (s *LoadCancelService) cancel(ctx context.Context, ref LoadJobRef, reason string, byAdmin bool) (LoadCancelResult, error) { + outcome, job, err := s.Registry.cancelLoadJob(ctx, ref, reason, byAdmin) + if err != nil { + return LoadCancelResult{}, err + } + switch outcome { + case CancelGone: + return LoadCancelResult{State: LoadCancelGone}, nil + case CancelConflict: + current := "" + if job != nil { + current = job.Generation + } + return LoadCancelResult{CurrentJobID: current}, fmt.Errorf("%w: %s", ErrLoadCancelConflict, current) + } + + // The attempt is failed and cancelled. Stop its remote work. + if !job.OpConfirmed && job.NodeID != "" && s.Stopper != nil { + addr := s.Registry.attemptAddress(ctx, job) + if StopLoadAttempt(ctx, s.Registry, s.Stopper, ref, job.NodeID, job.ReplicaIndex, addr, job.LegacyWorker) { + return LoadCancelResult{State: LoadCancelStopped}, nil + } + } + if job.OpConfirmed { + return LoadCancelResult{State: LoadCancelStopped}, nil + } + result := LoadCancelResult{State: LoadCancelStopping} + if fresh, gerr := s.Registry.GetLoadJob(ctx, ref.TrackingKey); gerr == nil && fresh != nil && fresh.StopDeadline != nil { + result.RetryAfter = max(time.Until(*fresh.StopDeadline), time.Second) + } + return result, nil +} + +// CancelModelOnNode cancels the load of modelName that runs on nodeID, or has +// no node yet. A load placed on another node is left alone: unloading a replica +// here must not cancel a load there. It reports the attempts it cancelled. +// +// An unload of a loaded replica calls this first and then carries on with the +// normal unload: stopping the load's operation never stops a model that +// finished loading, because the worker refuses a stop whose operation ended. +func (s *LoadCancelService) CancelModelOnNode(ctx context.Context, nodeID, modelName string) ([]LoadJobRef, error) { + job, err := s.Registry.GetLoadJob(ctx, modelName) + if err != nil || job == nil { + return nil, err + } + if job.NodeID != "" && job.NodeID != nodeID { + return nil, nil + } + if _, err := s.Cancel(ctx, job.Ref()); err != nil && !errors.Is(err, ErrLoadCancelConflict) { + return nil, err + } + return []LoadJobRef{job.Ref()}, nil +} + +// CancelNodeLoads cancels every load placed on nodeID. Deregistering or +// draining a node calls it before the node's rows are removed. +func (s *LoadCancelService) CancelNodeLoads(ctx context.Context, nodeID string) error { + jobs, err := s.Registry.ListLoadJobsOnNode(ctx, nodeID) + if err != nil { + return err + } + var errs []error + for _, job := range jobs { + if _, err := s.cancel(ctx, job.Ref(), "the worker was removed or is shutting down", false); err != nil && !errors.Is(err, ErrLoadCancelConflict) { + errs = append(errs, err) + } + } + return errors.Join(errs...) +} + +// retryLoadStops retries the stop of failed attempts whose remote work is not +// confirmed ended, once per reconciler pass, until the worker acknowledges or +// the stop deadline releases the job. +func (rc *ReplicaReconciler) retryLoadStops(ctx context.Context) { + var stopper LoadAttemptStopper + switch { + case rc.unloader != nil: + stopper = rc.unloader + case rc.adapter != nil: + stopper = rc.adapter + default: + return + } + jobs, err := rc.registry.ListLoadJobsAwaitingStop(ctx) + if err != nil { + xlog.Warn("Reconciler: failed to list load jobs awaiting stop", "error", err) + return + } + for _, job := range jobs { + StopLoadAttempt(ctx, rc.registry, stopper, job.Ref(), job.NodeID, job.ReplicaIndex, rc.registry.attemptAddress(ctx, &job), job.LegacyWorker) + } +} + +// loadAttemptRegistry is what StopLoadAttempt records its outcome on. +type loadAttemptRegistry interface { + ConfirmLoadOp(ctx context.Context, ref LoadJobRef) error + SetLegacyStopWindow(ctx context.Context, ref LoadJobRef, window time.Duration) error +} + +// StopLoadAttempt stops the remote work of one failed or cancelled attempt and +// records the outcome. It is the single place a load's work is stopped. +// +// A worker that names operations is sent a stop by operation id, with the +// address when known. A worker that does not (it reported no process instance) +// is sent a stop by exact process address and nothing else, and when the address +// is unknown no stop is claimed. An acknowledged stop confirms the attempt and +// shortens the hold. For a legacy worker that does not acknowledge, the hold is +// the load deadline, because nothing sooner bounds its work. It reports whether +// the stop was acknowledged. +func StopLoadAttempt(ctx context.Context, reg loadAttemptRegistry, stopper LoadAttemptStopper, ref LoadJobRef, nodeID string, replica int, addr string, legacy bool) bool { + acked := false + switch { + case legacy: + if addr != "" { + reply, err := stopper.StopModelReplica(ctx, nodeID, NodeModel{ModelName: ref.TrackingKey, ReplicaIndex: replica, Address: addr}, true) + acked = err == nil && reply.Error == "" && reply.Terminated + } + default: + acked = StopOperationAcked(ctx, stopper, nodeID, ref, replica, addr) + } + switch { + case acked: + if err := reg.ConfirmLoadOp(ctx, ref); err != nil && !errors.Is(err, ErrStaleLoadJob) { + xlog.Warn("Failed to record the stop confirmation", "model", ref.TrackingKey, "error", err) + } + case legacy: + if err := reg.SetLegacyStopWindow(ctx, ref, loadJobLegacyStopWindow); err != nil && !errors.Is(err, ErrStaleLoadJob) { + xlog.Warn("Failed to set the legacy stop window", "model", ref.TrackingKey, "error", err) + } + } + return acked +} + +// StopOperationAcked asks the worker to stop the operation of ref and reports +// whether it acknowledged: the process is gone, or was never there. An error, a +// refusal or silence is not an acknowledgement. +func StopOperationAcked(ctx context.Context, stopper LoadOperationStopper, nodeID string, ref LoadJobRef, replica int, addr string) bool { + reply, err := stopper.StopLoadOperation(ctx, nodeID, workerctl.ModelStopRequest{ + ModelName: ref.TrackingKey, + ProcessKey: model.BackendProcessKey(ref.TrackingKey, replica), + ExpectedAddress: addr, + OperationID: ref.Generation, + Force: true, + }) + if err != nil { + xlog.Warn("Stopping the load operation failed", "node", nodeID, "model", ref.TrackingKey, "error", err) + return false + } + if reply.Error != "" || !reply.Terminated { + xlog.Warn("The worker did not stop the load operation", "node", nodeID, "model", ref.TrackingKey, "error", reply.Error) + return false + } + return true +} + +// attemptAddress returns the backend address recorded on the attempt's replica +// row, or "" when the row is gone or never reached a backend. +func (r *NodeRegistry) attemptAddress(ctx context.Context, job *ModelLoadJob) string { + var nm NodeModel + if err := r.db.WithContext(ctx). + Where("node_id = ? AND model_name = ? AND replica_index = ? AND load_generation = ?", job.NodeID, job.TrackingKey, job.ReplicaIndex, job.Generation). + First(&nm).Error; err != nil { + return "" + } + return nm.Address +} diff --git a/core/services/nodes/load_job_generation_test.go b/core/services/nodes/load_job_generation_test.go new file mode 100644 index 000000000..b9a8e3f15 --- /dev/null +++ b/core/services/nodes/load_job_generation_test.go @@ -0,0 +1,461 @@ +// SPDX-License-Identifier: MIT +package nodes + +import ( + "context" + "errors" + "os" + "path/filepath" + "regexp" + "runtime" + "strings" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/protobuf/proto" + "gorm.io/gorm" + + "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// These specs run against a real database. They check what a caller can see +// (can the model load again, did the stale owner stop) and not which SQL ran. +var _ = Describe("Load job generation fencing", func() { + var ( + db *gorm.DB + registry *NodeRegistry + ctx context.Context + ) + + BeforeEach(func() { + if runtime.GOOS == "darwin" { + Skip("testcontainers requires Docker, not available on macOS CI") + } + db = testutil.SetupTestDB() + var err error + registry, err = NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + ctx = context.Background() + }) + + // release moves a failed job's stop deadline into the past, on the + // database clock. + release := func(trackingKey string) { + Expect(db.Exec("UPDATE model_load_jobs SET stop_deadline = now() - interval '1 second' WHERE tracking_key = ?", trackingKey).Error).To(Succeed()) + } + + loadJobCount := func(trackingKey string) int64 { + var n int64 + Expect(db.Model(&ModelLoadJob{}).Where("tracking_key = ?", trackingKey).Count(&n).Error).To(Succeed()) + return n + } + + Describe("stale owner writes", func() { + It("returns no rows for a heartbeat, fail and delete of a replaced attempt", func() { + a, claimed, err := registry.ClaimLoadJob(ctx, "aba-model", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(a.Generation).ToNot(BeEmpty()) + Expect(registry.DeleteLoadJob(ctx, a.Ref())).To(Succeed()) + + b, claimed, err := registry.ClaimLoadJob(ctx, "aba-model", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(b.Generation).ToNot(Equal(a.Generation)) + + for _, stale := range []LoadJobRef{a.Ref(), {TrackingKey: "aba-model"}} { + Expect(registry.UpdateLoadJob(ctx, stale, LoadJobUpdate{State: LoadJobStateLoading})).To(MatchError(ErrStaleLoadJob)) + Expect(registry.FailLoadJob(ctx, stale, "late failure", false)).To(MatchError(ErrStaleLoadJob)) + Expect(registry.DeleteLoadJob(ctx, stale)).To(MatchError(ErrStaleLoadJob)) + } + + current, err := registry.GetLoadJob(ctx, "aba-model") + Expect(err).ToNot(HaveOccurred()) + Expect(current.Ref()).To(Equal(b.Ref())) + Expect(current.State).To(Equal(LoadJobStatePending), "a stale write must not touch the current attempt") + }) + + It("does not report a database failure as a lost generation", func() { + job, _, err := registry.ClaimLoadJob(ctx, "cancelled-write", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + cancelled, cancel := context.WithCancel(ctx) + cancel() + writeErr := registry.UpdateLoadJob(cancelled, job.Ref(), LoadJobUpdate{}) + Expect(writeErr).To(HaveOccurred()) + Expect(writeErr).ToNot(MatchError(ErrStaleLoadJob)) + }) + + It("refuses a replica publish from a replaced attempt and from a context with no ownership", func() { + node := &BackendNode{Name: "worker-1", NodeType: NodeTypeBackend, Address: "10.0.0.1:50051"} + Expect(registry.Register(ctx, node, true)).To(Succeed()) + + a, _, err := registry.ClaimLoadJob(ctx, "publish-model", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.DeleteLoadJob(ctx, a.Ref())).To(Succeed()) + b, _, err := registry.ClaimLoadJob(ctx, "publish-model", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + + staleCtx := withLoadOwnership(ctx, a.Ref()) + Expect(registry.SetNodeModel(staleCtx, node.ID, "publish-model", 0, "staging", "10.0.0.1:9001", 0)).To(MatchError(ErrStaleLoadJob)) + Expect(registry.SetNodeModelLoadInfo(staleCtx, node.ID, "publish-model", 0, "llama-cpp", []byte("x"))).To(MatchError(ErrStaleLoadJob)) + Expect(registry.UpsertModelLoadInfo(staleCtx, "publish-model", "llama-cpp", []byte("x"))).To(MatchError(ErrStaleLoadJob)) + Expect(registry.RemoveNodeModel(staleCtx, node.ID, "publish-model", 0)).To(MatchError(ErrStaleLoadJob)) + + // The load path never runs without an ownership value. If it does, + // the write is refused instead of waved through. + unowned := withLoadPath(ctx) + Expect(registry.SetNodeModel(unowned, node.ID, "publish-model", 0, "staging", "10.0.0.1:9001", 0)).To(MatchError(ErrLoadOwnershipMissing)) + + var rows int64 + Expect(db.Model(&NodeModel{}).Where("model_name = ?", "publish-model").Count(&rows).Error).To(Succeed()) + Expect(rows).To(BeZero(), "no fenced write may leave a replica row behind") + + ownerCtx := withLoadOwnership(withLoadPath(ctx), b.Ref()) + Expect(registry.SetNodeModel(ownerCtx, node.ID, "publish-model", 0, "staging", "10.0.0.1:9001", 0)).To(Succeed()) + Expect(db.Model(&NodeModel{}).Where("model_name = ?", "publish-model").Count(&rows).Error).To(Succeed()) + Expect(rows).To(Equal(int64(1))) + + // Callers outside the load path are unaffected. + Expect(registry.SetNodeModel(ctx, node.ID, "other-model", 0, "loaded", "10.0.0.1:9002", 0)).To(Succeed()) + }) + }) + + Describe("failed jobs", func() { + It("frees a failed model without manual SQL, even after a frontend restart", func() { + failed, _, err := registry.ClaimLoadJob(ctx, "retry-model", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.FailLoadJob(ctx, failed.Ref(), "worker out of disk", false)).To(Succeed()) + + // Inside the grace window every caller sees the real cause. + got, claimed, err := registry.ClaimLoadJob(ctx, "retry-model", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(got.State).To(Equal(LoadJobStateFailed)) + Expect(got.LastError).To(Equal("worker out of disk")) + + release("retry-model") + + // A new registry stands in for a restarted frontend: the release + // must not depend on an in-process timer. + restarted, err := NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + fresh, claimed, err := restarted.ClaimLoadJob(ctx, "retry-model", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue(), "a failed job past its grace window must not block the model") + Expect(fresh.Generation).ToNot(Equal(failed.Generation)) + Expect(loadJobCount("retry-model")).To(Equal(int64(1))) + + // The old owner cannot touch the new attempt. + Expect(registry.UpdateLoadJob(ctx, failed.Ref(), LoadJobUpdate{State: LoadJobStateLoading})).To(MatchError(ErrStaleLoadJob)) + Expect(registry.DeleteLoadJob(ctx, failed.Ref())).To(MatchError(ErrStaleLoadJob)) + }) + + It("does not let the success path delete a failed job", func() { + failed, _, err := registry.ClaimLoadJob(ctx, "failed-delete", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.FailLoadJob(ctx, failed.Ref(), "boom", false)).To(Succeed()) + Expect(registry.DeleteLoadJob(ctx, failed.Ref())).To(MatchError(ErrStaleLoadJob)) + Expect(loadJobCount("failed-delete")).To(Equal(int64(1))) + + // The failed row is removed by its own, grace-gated call. + Expect(registry.DeleteFailedLoadJob(ctx, failed.Ref())).To(MatchError(ErrStaleLoadJob)) + release("failed-delete") + Expect(registry.DeleteFailedLoadJob(ctx, failed.Ref())).To(Succeed()) + Expect(loadJobCount("failed-delete")).To(BeZero()) + }) + }) + + Describe("legacy rows", func() { + // An old binary created the table without a generation column. A row + // there has a NULL generation once the column exists. + legacySchema := func() *gorm.DB { + legacyDB := testutil.SetupTestDB() + Expect(legacyDB.Exec(`CREATE TABLE model_load_jobs ( + tracking_key varchar(255) PRIMARY KEY, state varchar(16) NOT NULL, owner_replica varchar(64), + node_id varchar(36), node_name varchar(255), replica_index bigint, bytes_sent bigint, + total_bytes bigint, file_index bigint, total_files bigint, last_error text, + started_at timestamptz, created_at timestamptz, updated_at timestamptz, last_progress timestamptz)`).Error).To(Succeed()) + return legacyDB + } + insertLegacy := func(legacyDB *gorm.DB, key string, age time.Duration) { + Expect(legacyDB.Exec(`INSERT INTO model_load_jobs (tracking_key, state, owner_replica, last_progress, created_at, updated_at) + VALUES (?, 'staging', 'old-frontend', ?, now(), now())`, key, time.Now().Add(-age)).Error).To(Succeed()) + } + + It("gives a pre-existing row a generation and an expired lease", func() { + legacyDB := legacySchema() + insertLegacy(legacyDB, "legacy-dead", 5*time.Minute) + insertLegacy(legacyDB, "legacy-live", time.Second) + + migrated, err := NewNodeRegistry(legacyDB) + Expect(err).ToNot(HaveOccurred()) + + var jobs []ModelLoadJob + Expect(legacyDB.Find(&jobs).Error).To(Succeed()) + Expect(jobs).To(HaveLen(2)) + for _, j := range jobs { + Expect(j.Generation).ToNot(BeEmpty(), "the migration must give every legacy row a generation") + } + + // No old binary renews a lease, so every legacy row is expired, + // live or not. It is released after the stop window. + for _, key := range []string{"legacy-dead", "legacy-live"} { + held, claimed, err := migrated.ClaimLoadJob(ctx, key, "new-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(held.State).To(Equal(LoadJobStateFailed)) + } + Expect(legacyDB.Exec("UPDATE model_load_jobs SET stop_deadline = now() - interval '1 second'").Error).To(Succeed()) + dead, claimed, err := migrated.ClaimLoadJob(ctx, "legacy-dead", "new-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(dead.Generation).ToNot(BeEmpty()) + }) + + It("treats a row written later by an old binary (empty generation) the same way", func() { + legacyDB := legacySchema() + migrated, err := NewNodeRegistry(legacyDB) + Expect(err).ToNot(HaveOccurred()) + insertLegacy(legacyDB, "late-legacy", 5*time.Minute) + + var row ModelLoadJob + Expect(legacyDB.First(&row, "tracking_key = ?", "late-legacy").Error).To(Succeed()) + Expect(row.Generation).To(BeEmpty()) + + // No new owner can hold an empty generation, so no write may match it. + Expect(migrated.UpdateLoadJob(ctx, row.Ref(), LoadJobUpdate{})).To(MatchError(ErrStaleLoadJob)) + Expect(migrated.FailLoadJob(ctx, row.Ref(), "x", false)).To(MatchError(ErrStaleLoadJob)) + Expect(migrated.DeleteLoadJob(ctx, row.Ref())).To(MatchError(ErrStaleLoadJob)) + + held, claimed, err := migrated.ClaimLoadJob(ctx, "late-legacy", "new-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(held.State).To(Equal(LoadJobStateFailed)) + Expect(legacyDB.Exec("UPDATE model_load_jobs SET stop_deadline = now() - interval '1 second'").Error).To(Succeed()) + job, claimed, err := migrated.ClaimLoadJob(ctx, "late-legacy", "new-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(job.Generation).ToNot(BeEmpty()) + }) + }) + + Describe("the owner loop", func() { + var ( + router *SmartRouter + unloader *fakeUnloader + backend *stubBackend + node *BackendNode + ) + + BeforeEach(func() { + node = &BackendNode{ + Name: "worker-1", NodeType: NodeTypeBackend, Address: "10.0.0.1:50051", + TotalVRAM: 64_000_000_000, AvailableVRAM: 64_000_000_000, + } + Expect(registry.Register(ctx, node, true)).To(Succeed()) + backend = &stubBackend{healthResult: true, loadResult: &pb.Result{Success: true}} + unloader = &fakeUnloader{installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}} + router = NewSmartRouter(registry, SmartRouterOptions{ + Unloader: unloader, ClientFactory: &stubClientFactory{client: backend}, DB: db, + }) + }) + + replaceJob := func(key string) LoadJobRef { + Expect(db.Where("tracking_key = ?", key).Delete(&ModelLoadJob{}).Error).To(Succeed()) + next, claimed, err := registry.ClaimLoadJob(ctx, key, "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + return next.Ref() + } + + It("cancels the work when another generation took the model, and leaves that generation alone", func() { + job, _, err := registry.ClaimLoadJob(ctx, "stale-owner", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + + workCtx := make(chan context.Context, 1) + result := make(chan error, 1) + go func() { + defer GinkgoRecover() + result <- router.runLoadOwner(ctx, job.Ref(), func(c context.Context) error { + workCtx <- c + <-c.Done() + return c.Err() + }) + }() + var running context.Context + Eventually(workCtx).Should(Receive(&running)) + + next := replaceJob("stale-owner") + + Eventually(running.Done(), 10*time.Second).Should(BeClosed(), "the stale owner must stop its own work") + var ownerErr error + Eventually(result, 10*time.Second).Should(Receive(&ownerErr)) + Expect(errors.Is(ownerErr, ErrStaleLoadJob)).To(BeTrue(), "the stale error must reach the caller, not be swallowed") + + current, err := registry.GetLoadJob(ctx, "stale-owner") + Expect(err).ToNot(HaveOccurred()) + Expect(current).ToNot(BeNil()) + Expect(current.Ref()).To(Equal(next)) + Expect(current.State).To(Equal(LoadJobStatePending), "the stale owner must not fail or delete the new attempt") + }) + + It("deletes the job on success and records the cause on failure", func() { + ok, _, err := registry.ClaimLoadJob(ctx, "owner-ok", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(router.runLoadOwner(ctx, ok.Ref(), func(context.Context) error { return nil })).To(Succeed()) + Expect(loadJobCount("owner-ok")).To(BeZero()) + + bad, _, err := registry.ClaimLoadJob(ctx, "owner-bad", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + err = router.runLoadOwner(ctx, bad.Ref(), func(context.Context) error { return errors.New("remote load failed") }) + Expect(err).To(MatchError(ContainSubstring("remote load failed"))) + row, err := registry.GetLoadJob(ctx, "owner-bad") + Expect(err).ToNot(HaveOccurred()) + Expect(row.State).To(Equal(LoadJobStateFailed)) + Expect(row.LastError).To(ContainSubstring("remote load failed")) + }) + + It("wakes only the waiters of the generation that finished", func() { + a, _, err := registry.ClaimLoadJob(ctx, "waiters", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.DeleteLoadJob(ctx, a.Ref())).To(Succeed()) + b, _, err := registry.ClaimLoadJob(ctx, "waiters", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + + waiter := router.loadWaiterChan(loadWaiterKey(b.Ref())) + Expect(router.finishLoadJob(ctx, a.Ref())).To(MatchError(ErrStaleLoadJob)) + Consistently(waiter, 200*time.Millisecond).ShouldNot(BeClosed(), "a delayed finish of an old generation must not wake the new one") + Expect(router.finishLoadJob(ctx, b.Ref())).To(Succeed()) + Eventually(waiter).Should(BeClosed()) + }) + + It("releases a waiter registration when the waiter leaves early", func() { + job, _, err := registry.ClaimLoadJob(ctx, "leaving", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + waitCtx, cancel := context.WithTimeout(ctx, 100*time.Millisecond) + defer cancel() + waiter := router.loadWaiterChan(loadWaiterKey(job.Ref())) + Expect(router.waitForLoadJob(waitCtx, job.Ref(), waiter)).To(HaveOccurred()) + router.loadWaitersMu.Lock() + defer router.loadWaitersMu.Unlock() + Expect(router.loadWaiters).ToNot(HaveKey(loadWaiterKey(job.Ref()))) + }) + + It("retries a model whose first load failed, with no SQL in between", func() { + // The remote load fails with an error the backend answered. + backend.loadResult = &pb.Result{Success: false, Message: "unsupported architecture"} + opts := &pb.ModelOptions{Model: "models/retry.gguf"} + _, err := router.Route(ctx, "retry-route", "models/retry.gguf", "llama-cpp", "", opts, false) + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("unsupported architecture")) + + var failed ModelLoadJob + Eventually(func() string { + _ = db.First(&failed, "tracking_key = ?", "retry-route").Error + return failed.State + }, 5*time.Second, 50*time.Millisecond).Should(Equal(LoadJobStateFailed)) + + // Inside the grace window the caller gets the real cause quickly. + _, err = router.Route(ctx, "retry-route", "models/retry.gguf", "llama-cpp", "", opts, false) + Expect(err).To(MatchError(ContainSubstring("unsupported architecture"))) + + // The backend is fixed. The grace window ends. No manual cleanup. + backend.mu.Lock() + backend.loadResult = &pb.Result{Success: true} + backend.mu.Unlock() + release("retry-route") + + res, err := router.Route(ctx, "retry-route", "models/retry.gguf", "llama-cpp", "", opts, false) + Expect(err).ToNot(HaveOccurred()) + res.Release() + Eventually(func() int64 { return loadJobCount("retry-route") }, 5*time.Second, 50*time.Millisecond).Should(BeZero()) + }) + + It("runs the reconciler path under the same owner loop", func() { + const modelName = "reconciled" + blob, err := proto.Marshal(&pb.ModelOptions{Model: "models/reconciled.gguf"}) + Expect(err).ToNot(HaveOccurred()) + Expect(registry.UpsertModelLoadInfo(ctx, modelName, "llama-cpp", blob)).To(Succeed()) + + release := make(chan struct{}) + started := make(chan struct{}) + unloader.installHook = func() { close(started); <-release } + defer close(release) + + result := make(chan error, 1) + go func() { + defer GinkgoRecover() + _, err := router.ScheduleAndLoadModel(ctx, modelName, nil) + result <- err + }() + Eventually(started, 10*time.Second).Should(BeClosed()) + + // The reconciler load is a durable job, so a request for the same + // model waits for it instead of scheduling a second copy. + held, err := registry.GetLoadJob(ctx, modelName) + Expect(err).ToNot(HaveOccurred()) + Expect(held).ToNot(BeNil()) + _, claimed, err := registry.ClaimLoadJob(ctx, modelName, "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + + next := replaceJob(modelName) + + var loadErr error + Eventually(result, 10*time.Second).Should(Receive(&loadErr)) + Expect(errors.Is(loadErr, ErrStaleLoadJob)).To(BeTrue(), "the reconciler path must stop when it loses the generation") + + current, err := registry.GetLoadJob(ctx, modelName) + Expect(err).ToNot(HaveOccurred()) + Expect(current.Ref()).To(Equal(next)) + Expect(current.State).To(Equal(LoadJobStatePending)) + }) + + It("does not start a reconciler load while another owner holds the model", func() { + const modelName = "held-model" + blob, err := proto.Marshal(&pb.ModelOptions{Model: "models/held.gguf"}) + Expect(err).ToNot(HaveOccurred()) + Expect(registry.UpsertModelLoadInfo(ctx, modelName, "llama-cpp", blob)).To(Succeed()) + _, claimed, err := registry.ClaimLoadJob(ctx, modelName, "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + + _, err = router.ScheduleAndLoadModel(ctx, modelName, nil) + Expect(err).To(HaveOccurred()) + unloader.mu.Lock() + defer unloader.mu.Unlock() + Expect(unloader.installCalls).To(BeEmpty()) + }) + }) +}) + +// Every write of a durable load job must carry the generation predicate, so +// the predicate lives in one helper. This guards against a new write that +// bypasses it. +var _ = Describe("Load job write discipline", func() { + It("keeps every ModelLoadJob write inside model_load_job.go", func() { + write := regexp.MustCompile(`\.(Model|Delete|Update|Updates|UpdateColumn|Save|Exec)\((&ModelLoadJob\{\}|"[^"]*model_load_jobs)|Table\("model_load_jobs"\)`) + files, err := filepath.Glob("*.go") + Expect(err).ToNot(HaveOccurred()) + for _, f := range files { + if strings.HasSuffix(f, "_test.go") || f == "model_load_job.go" { + continue + } + src, err := os.ReadFile(f) + Expect(err).ToNot(HaveOccurred()) + Expect(write.Match(src)).To(BeFalse(), "%s writes model_load_jobs outside the ownedLoadJob helper", f) + } + }) + + It("builds every write in model_load_job.go from ownedLoadJob", func() { + src, err := os.ReadFile("model_load_job.go") + Expect(err).ToNot(HaveOccurred()) + // The helper is the only place a query on the table starts. + Expect(strings.Count(string(src), "Model(&ModelLoadJob{})")).To(Equal(1)) + // A delete with a free-form condition would skip the predicate. + Expect(string(src)).ToNot(MatchRegexp(`Delete\(&ModelLoadJob\{\},`)) + }) +}) diff --git a/core/services/nodes/load_job_lease_test.go b/core/services/nodes/load_job_lease_test.go new file mode 100644 index 000000000..290462ed8 --- /dev/null +++ b/core/services/nodes/load_job_lease_test.go @@ -0,0 +1,349 @@ +// SPDX-License-Identifier: MIT +package nodes + +import ( + "context" + "errors" + "runtime" + "sync/atomic" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gorm.io/gorm" + + "github.com/mudler/LocalAI/core/services/testutil" +) + +// flakyJobRegistry fails lease renewals on demand, the way a database that is +// briefly unreachable would. +type flakyJobRegistry struct { + *NodeRegistry + failRenewals atomic.Bool +} + +func (f *flakyJobRegistry) UpdateLoadJob(ctx context.Context, ref LoadJobRef, u LoadJobUpdate) error { + if f.failRenewals.Load() { + return errors.New("database unreachable") + } + return f.NodeRegistry.UpdateLoadJob(ctx, ref, u) +} + +// These specs move time by editing the stored deadlines relative to the +// database clock, because the lease is decided by the database and never by the +// frontend clock. +var _ = Describe("Load job lease and reclaim", func() { + var ( + db *gorm.DB + registry *NodeRegistry + ctx context.Context + ) + + BeforeEach(func() { + if runtime.GOOS == "darwin" { + Skip("testcontainers requires Docker, not available on macOS CI") + } + db = testutil.SetupTestDB() + var err error + registry, err = NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + ctx = context.Background() + }) + + // shift sets a deadline column to the database's now() plus secs. + shift := func(key, column string, secs int) { + Expect(db.Exec("UPDATE model_load_jobs SET "+column+" = now() + make_interval(secs => ?) WHERE tracking_key = ?", secs, key).Error).To(Succeed()) + } + // dbSecondsUntil reads how far a deadline is ahead of the database clock. + dbSecondsUntil := func(key, column string) float64 { + var secs float64 + Expect(db.Raw("SELECT EXTRACT(EPOCH FROM ("+column+" - now())) FROM model_load_jobs WHERE tracking_key = ?", key).Scan(&secs).Error).To(Succeed()) + return secs + } + loadJob := func(key string) *ModelLoadJob { + job, err := registry.GetLoadJob(ctx, key) + Expect(err).ToNot(HaveOccurred()) + return job + } + + Describe("an owner that dies", func() { + It("is released by the reconciler alone, and the next request then loads", func() { + node := &BackendNode{Name: "n1", NodeType: NodeTypeBackend, Address: "10.0.0.1:50051"} + Expect(registry.Register(ctx, node, true)).To(Succeed()) + rc := NewReplicaReconciler(ReplicaReconcilerOptions{Registry: registry, DB: db}) + + dead, claimed, err := registry.ClaimLoadJob(ctx, "crashed", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(dbSecondsUntil("crashed", "lease_until")).To(BeNumerically("~", loadJobLeaseTTL.Seconds(), 3)) + // The dead owner had published a staging replica before it died. + Expect(registry.SetNodeModel(withLoadOwnership(ctx, dead.Ref()), node.ID, "crashed", 0, "staging", "10.0.0.1:9001", 0)).To(Succeed()) + + // While the lease is live nobody may take the model, and the sweep + // leaves everything alone. + _, claimed, err = registry.ClaimLoadJob(ctx, "crashed", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + rc.reclaimAbandonedLoads(ctx) + Expect(loadJob("crashed").State).To(Equal(LoadJobStatePending)) + + // The owner never renews. The lease runs out; no request arrives. + shift("crashed", "lease_until", -1) + rc.reclaimAbandonedLoads(ctx) + failed := loadJob("crashed") + Expect(failed).ToNot(BeNil()) + Expect(failed.State).To(Equal(LoadJobStateFailed)) + Expect(failed.LastError).To(ContainSubstring("lease")) + Expect(dbSecondsUntil("crashed", "stop_deadline")).To(BeNumerically("~", loadJobStopWindow.Seconds(), 3), + "the slot is held for the stop window, not for ever") + + // A request inside the stop window is told the cause and is not + // started a second load on top of work that may still run. + held, claimed, err := registry.ClaimLoadJob(ctx, "crashed", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(held.State).To(Equal(LoadJobStateFailed)) + var rows int64 + Expect(db.Model(&NodeModel{}).Where("model_name = ?", "crashed").Count(&rows).Error).To(Succeed()) + Expect(rows).To(Equal(int64(1)), "the replica slot stays reserved until the job is released") + + // The stop window passes. The sweep releases the job and its replica. + shift("crashed", "stop_deadline", -1) + rc.reclaimAbandonedLoads(ctx) + Expect(loadJob("crashed")).To(BeNil()) + Expect(db.Model(&NodeModel{}).Where("model_name = ?", "crashed").Count(&rows).Error).To(Succeed()) + Expect(rows).To(BeZero()) + + next, claimed, err := registry.ClaimLoadJob(ctx, "crashed", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(next.Generation).ToNot(Equal(dead.Generation)) + }) + + It("is reclaimed by a request alone when no reconciler runs", func() { + dead, _, err := registry.ClaimLoadJob(ctx, "lazy", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + shift("lazy", "lease_until", -1) + + first, claimed, err := registry.ClaimLoadJob(ctx, "lazy", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(first.State).To(Equal(LoadJobStateFailed)) + + shift("lazy", "stop_deadline", -1) + next, claimed, err := registry.ClaimLoadJob(ctx, "lazy", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(next.Generation).ToNot(Equal(dead.Generation)) + Expect(registry.UpdateLoadJob(ctx, dead.Ref(), LoadJobUpdate{})).To(MatchError(ErrStaleLoadJob)) + }) + }) + + Describe("clock skew", func() { + It("lets the database clock decide the lease, not the frontend clock", func() { + for _, skew := range []time.Duration{10 * time.Minute, -10 * time.Minute} { + key := "skew-" + skew.String() + registry.clock = func() time.Time { return time.Now().Add(skew) } + + live, claimed, err := registry.ClaimLoadJob(ctx, key, "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + + // A live lease is not expired by a frontend clock that runs ahead. + other, claimed, err := registry.ClaimLoadJob(ctx, key, "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(other.State).ToNot(Equal(LoadJobStateFailed)) + + // Renewal extends the lease by the database's time. + Expect(registry.UpdateLoadJob(ctx, live.Ref(), LoadJobUpdate{})).To(Succeed()) + Expect(dbSecondsUntil(key, "lease_until")).To(BeNumerically("~", loadJobLeaseTTL.Seconds(), 3)) + + // A dead lease is not kept alive by a frontend clock that runs behind. + shift(key, "lease_until", -1) + expired, claimed, err := registry.ClaimLoadJob(ctx, key, "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(expired.State).To(Equal(LoadJobStateFailed)) + } + }) + }) + + Describe("lease renewal in the owner loop", func() { + var ( + flaky *flakyJobRegistry + router *SmartRouter + ) + BeforeEach(func() { + flaky = &flakyJobRegistry{NodeRegistry: registry} + registry.leaseTTL = 4 * time.Second + router = NewSmartRouter(flaky, SmartRouterOptions{DB: db}) + router.leaseTTL = registry.leaseTTL + }) + + It("renews the lease with every heartbeat", func() { + job, _, err := registry.ClaimLoadJob(ctx, "renewing", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + shift("renewing", "lease_until", 1) + + release := make(chan struct{}) + done := make(chan error, 1) + go func() { + defer GinkgoRecover() + done <- router.runLoadOwner(ctx, job.Ref(), func(context.Context) error { <-release; return nil }) + }() + Eventually(func() float64 { return dbSecondsUntil("renewing", "lease_until") }, 5*time.Second, 200*time.Millisecond). + Should(BeNumerically(">", 2), "a live owner must push its lease forward") + close(release) + Eventually(done, 5*time.Second).Should(Receive(BeNil())) + }) + + It("keeps working through a short database failure", func() { + job, _, err := registry.ClaimLoadJob(ctx, "blip", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + + started := make(chan context.Context, 1) + release := make(chan struct{}) + go func() { + defer GinkgoRecover() + _ = router.runLoadOwner(ctx, job.Ref(), func(c context.Context) error { + started <- c + select { + case <-c.Done(): + return c.Err() + case <-release: + return nil + } + }) + }() + var workCtx context.Context + Eventually(started).Should(Receive(&workCtx)) + + flaky.failRenewals.Store(true) + time.Sleep(2 * time.Second) // under the TTL + flaky.failRenewals.Store(false) + Consistently(workCtx.Done(), 2*time.Second).ShouldNot(BeClosed()) + close(release) + }) + + It("stops its own work when it can no longer extend the lease", func() { + job, _, err := registry.ClaimLoadJob(ctx, "cut-off", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + + flaky.failRenewals.Store(true) + result := make(chan error, 1) + go func() { + defer GinkgoRecover() + result <- router.runLoadOwner(ctx, job.Ref(), func(c context.Context) error { + <-c.Done() + return c.Err() + }) + }() + + var ownerErr error + Eventually(result, 15*time.Second).Should(Receive(&ownerErr)) + Expect(errors.Is(ownerErr, ErrLoadLeaseExpired)).To(BeTrue(), "an owner must not outlive a lease it cannot extend") + // The database came back, so the owner could record the failure. + Expect(loadJob("cut-off").State).To(Equal(LoadJobStateFailed)) + }) + }) + + Describe("failed jobs", func() { + It("clears a failure the backend answered after the report window, with no SQL", func() { + job, _, err := registry.ClaimLoadJob(ctx, "answered", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.FailLoadJob(ctx, job.Ref(), "unsupported architecture", false)).To(Succeed()) + Expect(dbSecondsUntil("answered", "stop_deadline")).To(BeNumerically("~", loadJobFailureReport.Seconds(), 3)) + + _, claimed, err := registry.ClaimLoadJob(ctx, "answered", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + + shift("answered", "stop_deadline", -1) + restarted, err := NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + next, claimed, err := restarted.ClaimLoadJob(ctx, "answered", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(next.Generation).ToNot(Equal(job.Generation)) + }) + + It("holds the slot for the stop window when remote work may still run", func() { + job, _, err := registry.ClaimLoadJob(ctx, "maybe-running", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.FailLoadJob(ctx, job.Ref(), "context deadline exceeded", true)).To(Succeed()) + Expect(dbSecondsUntil("maybe-running", "stop_deadline")).To(BeNumerically("~", loadJobStopWindow.Seconds(), 3)) + }) + }) + + Describe("legacy rows", func() { + It("treats a row from before the lease as expired and releases it after the stop window", func() { + legacyDB := testutil.SetupTestDB() + Expect(legacyDB.Exec(`CREATE TABLE model_load_jobs ( + tracking_key varchar(255) PRIMARY KEY, state varchar(16) NOT NULL, owner_replica varchar(64), + node_id varchar(36), node_name varchar(255), replica_index bigint, bytes_sent bigint, + total_bytes bigint, file_index bigint, total_files bigint, last_error text, + started_at timestamptz, created_at timestamptz, updated_at timestamptz, last_progress timestamptz)`).Error).To(Succeed()) + Expect(legacyDB.Exec(`INSERT INTO model_load_jobs (tracking_key, state, owner_replica, last_progress, created_at, updated_at) + VALUES ('old-load', 'staging', 'old-frontend', now(), now(), now())`).Error).To(Succeed()) + + migrated, err := NewNodeRegistry(legacyDB) + Expect(err).ToNot(HaveOccurred()) + + held, claimed, err := migrated.ClaimLoadJob(ctx, "old-load", "new-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(held.State).To(Equal(LoadJobStateFailed), "an old binary never renews a lease, so its row is expired") + + Expect(legacyDB.Exec("UPDATE model_load_jobs SET stop_deadline = now() - interval '1 second'").Error).To(Succeed()) + next, claimed, err := migrated.ClaimLoadJob(ctx, "old-load", "new-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(next.Generation).ToNot(BeEmpty()) + }) + + It("gives a failed row from an old binary a stop window instead of releasing it at once", func() { + Expect(db.Exec(`INSERT INTO model_load_jobs (tracking_key, state, owner_replica, last_error, last_progress, created_at, updated_at) + VALUES ('old-failed', 'failed', 'old-frontend', 'boom', now(), now(), now())`).Error).To(Succeed()) + held, claimed, err := registry.ClaimLoadJob(ctx, "old-failed", "new-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(held.LastError).To(Equal("boom")) + Expect(dbSecondsUntil("old-failed", "stop_deadline")).To(BeNumerically("~", loadJobStopWindow.Seconds(), 3)) + }) + }) + + Describe("replicas of an abandoned attempt", func() { + It("reaps by generation and leaves the current attempt's replica alone", func() { + a := &BackendNode{Name: "n-a", NodeType: NodeTypeBackend, Address: "10.0.0.1:50051"} + b := &BackendNode{Name: "n-b", NodeType: NodeTypeBackend, Address: "10.0.0.2:50051"} + Expect(registry.Register(ctx, a, true)).To(Succeed()) + Expect(registry.Register(ctx, b, true)).To(Succeed()) + rc := NewReplicaReconciler(ReplicaReconcilerOptions{Registry: registry, DB: db}) + + old, _, err := registry.ClaimLoadJob(ctx, "generations", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.SetNodeModel(withLoadOwnership(ctx, old.Ref()), a.ID, "generations", 0, "staging", "x:1", 0)).To(Succeed()) + // The owner dies, the job is released and another attempt takes over. + shift("generations", "lease_until", -1) + rc.reclaimAbandonedLoads(ctx) + shift("generations", "stop_deadline", -1) + rc.reclaimAbandonedLoads(ctx) + current, claimed, err := registry.ClaimLoadJob(ctx, "generations", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(registry.SetNodeModel(withLoadOwnership(ctx, current.Ref()), b.ID, "generations", 0, "staging", "y:1", 0)).To(Succeed()) + // A row of the dead attempt that survived (a late write before the + // release) belongs to a generation that no longer has a job. + Expect(db.Create(&NodeModel{ID: "straggler", NodeID: a.ID, ModelName: "generations", ReplicaIndex: 1, + State: "staging", LoadGeneration: old.Generation}).Error).To(Succeed()) + + rc.reclaimAbandonedLoads(ctx) + + var left []NodeModel + Expect(db.Where("model_name = ?", "generations").Find(&left).Error).To(Succeed()) + Expect(left).To(HaveLen(1)) + Expect(left[0].LoadGeneration).To(Equal(current.Generation)) + }) + }) +}) diff --git a/core/services/nodes/load_job_phase.go b/core/services/nodes/load_job_phase.go index ece0d099d..59f4be436 100644 --- a/core/services/nodes/load_job_phase.go +++ b/core/services/nodes/load_job_phase.go @@ -19,6 +19,10 @@ type loadPhaseReporter struct { nodeID string nodeName string replicaIndex int + // address is the backend's gRPC address, known once the install replied. + address string + // legacyWorker is set when the node's worker does not track operations. + legacyWorker bool } type loadPhaseKey struct{} @@ -41,6 +45,7 @@ func (p *loadPhaseReporter) snapshot() LoadJobUpdate { NodeID: p.nodeID, NodeName: p.nodeName, ReplicaIndex: p.replicaIndex, + LegacyWorker: p.legacyWorker, } } @@ -62,3 +67,35 @@ func reportLoadPhase(ctx context.Context, state string, node *BackendNode, repli p.set(state, node, replicaIndex) } } + +// placement returns where the load runs, once a node was chosen. +func (p *loadPhaseReporter) placement() (nodeID string, replicaIndex int, legacyWorker bool) { + p.mu.Lock() + defer p.mu.Unlock() + return p.nodeID, p.replicaIndex, p.legacyWorker +} + +func (p *loadPhaseReporter) backendAddress() string { + p.mu.Lock() + defer p.mu.Unlock() + return p.address +} + +// reportLoadAddress records the backend address of the load, for a stop of a +// worker that cannot name operations. +func reportLoadAddress(ctx context.Context, addr string) { + if p, ok := ctx.Value(loadPhaseKey{}).(*loadPhaseReporter); ok { + p.mu.Lock() + p.address = addr + p.mu.Unlock() + } +} + +// markLegacyWorker records that the load's worker cannot confirm a stop. +func markLegacyWorker(ctx context.Context) { + if p, ok := ctx.Value(loadPhaseKey{}).(*loadPhaseReporter); ok { + p.mu.Lock() + p.legacyWorker = true + p.mu.Unlock() + } +} diff --git a/core/services/nodes/load_job_runner.go b/core/services/nodes/load_job_runner.go index bfb18b1e5..7c7d2cab8 100644 --- a/core/services/nodes/load_job_runner.go +++ b/core/services/nodes/load_job_runner.go @@ -2,10 +2,13 @@ package nodes import ( "context" + "errors" "fmt" + "sync/atomic" "time" "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/xlog" ) @@ -38,19 +41,13 @@ func (r *SmartRouter) routeViaLoadJob(ctx context.Context, att *routeAttempt) (* } for range maxColdLoadRounds { - // Register interest BEFORE claiming, so a job that finishes immediately - // cannot close the channel before this waiter exists. - waiter := r.loadWaiterChan(att.trackingKey) - job, claimed, err := r.registry.ClaimLoadJob(ctx, att.trackingKey, ReplicaID()) if err != nil { - // A broken job table must not make the model unroutable: fall back - // to loading inline, which is what every release before this did. - xlog.Warn("Claiming the model load job failed; loading inline instead", - "model", att.trackingKey, "error", err) - loadCtx, cancelLoad := r.newColdLoadContext(context.WithoutCancel(ctx)) - defer cancelLoad() - return r.coldLoad(loadCtx, att, 1) + // Without the job row nothing fences this load, so a load without + // it would publish replicas nobody can cancel. Fail the request. + // Models that are already loaded keep serving: the warm path + // does not read the job table. + return nil, fmt.Errorf("claiming the load of model %s: %w", att.trackingKey, err) } switch { @@ -60,20 +57,25 @@ func (r *SmartRouter) routeViaLoadJob(ctx context.Context, att *routeAttempt) (* // the lock. Without it the claim would schedule a second copy of a // model that is already up. if result := r.tryWarmPath(ctx, att); result != nil { - r.finishLoadJob(ctx, att.trackingKey) + _ = r.finishLoadJob(ctx, job.Ref()) // already logged; the warm result stands return result, nil } - r.startLoadJob(ctx, att) + r.startLoadJob(ctx, att, job.Ref()) case job != nil && job.State == LoadJobStateFailed: // Inside the failure grace window: report the real cause rather // than silently starting a fresh load of a model that just failed. - return nil, fmt.Errorf("loading model %s: %s", att.trackingKey, job.LastError) + return nil, NewLoadHeldError(job) default: xlog.Info("Model is already loading on another replica; waiting for it", "model", att.trackingKey, "state", job.State, "node", job.NodeName, "owner", job.OwnerReplica) } - if err := r.waitForLoadJob(waitCtx, att.trackingKey, waiter); err != nil { + // The waiter is keyed by the generation it waits for, so the end of an + // older attempt can never wake it. A job that ended before this + // registration is caught by the authority check inside waitForLoadJob. + ref := job.Ref() + waiter := r.loadWaiterChan(loadWaiterKey(ref)) + if err := r.waitForLoadJob(waitCtx, ref, waiter); err != nil { // The caller's own context is still live, so it was the wait budget // that ran out, not the client giving up: answer with progress. if ctx.Err() == nil && waitCtx.Err() != nil { @@ -118,7 +120,7 @@ func (r *SmartRouter) loadingAnswer(ctx context.Context, trackingKey string, bud return fmt.Errorf("timed out waiting for model %s to load", trackingKey) } if job.State == LoadJobStateFailed { - return fmt.Errorf("loading model %s: %s", trackingKey, job.LastError) + return NewLoadHeldError(job) } return newModelLoadingError(job, budget) } @@ -127,61 +129,167 @@ func (r *SmartRouter) loadingAnswer(ctx context.Context, trackingKey string, bud // request that triggered it. The job is owned by its record, not by that // request: the client may disconnect, be retried onto another replica, or time // out, and the transfer keeps going. -func (r *SmartRouter) startLoadJob(ctx context.Context, att *routeAttempt) { - trackingKey := att.trackingKey +func (r *SmartRouter) startLoadJob(ctx context.Context, att *routeAttempt, ref LoadJobRef) { // Keep the request's context VALUES (prefix chain and friends) but none of // its cancellation — see newColdLoadContext. parent := context.WithoutCancel(ctx) go func() { - loadCtx, cancelLoad := r.newColdLoadContext(parent) - defer cancelLoad() + // runLoadOwner books the outcome and logs it; there is no caller left + // to return the error to. + _ = r.runLoadOwner(parent, ref, func(ownerCtx context.Context) error { + loadCtx, cancelLoad := r.newColdLoadContext(ownerCtx) + defer cancelLoad() + _, err := r.coldLoad(loadCtx, att, 0) + return err + }) + }() +} - phase := newLoadPhaseReporter() - loadCtx = withLoadPhaseReporter(loadCtx, phase) +// runLoadOwner is the one loop that owns a claimed load, for the request path +// and the reconciler path alike. It runs work under the job's generation, +// heartbeats the row, and ends with a conditional fail or delete. +// +// The heartbeat doubles as the ownership check: when it finds the job gone or +// held by another generation, the work context is cancelled. An owner that lost +// its job stops, and its late writes find zero rows. The stale error is +// returned to the caller instead of being swallowed, so the caller can tell a +// lost job from a failed load. +func (r *SmartRouter) runLoadOwner(ctx context.Context, ref LoadJobRef, work func(context.Context) error) error { + ownerCtx, cancel := context.WithCancelCause(withLoadOwnership(ctx, ref)) + defer cancel(nil) - stopHeartbeat := r.startLoadJobHeartbeat(parent, trackingKey, phase) + phase := newLoadPhaseReporter() + ownerCtx = withLoadPhaseReporter(ownerCtx, phase) - _, err := r.coldLoad(loadCtx, att, 0) + stopHeartbeat := r.startLoadJobHeartbeat(ownerCtx, ref, phase, cancel) + err := work(ownerCtx) + stopHeartbeat() + nodeID, replica, legacy := phase.placement() + // Record the placement now. The heartbeat writes it once a second, and a + // failure that comes sooner would leave the job with no node, so no stop + // could find the work. + if nodeID != "" { + flushCtx, cancelFlush := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second) + if ferr := r.registry.UpdateLoadJob(flushCtx, ref, phase.snapshot()); ferr != nil && !errors.Is(ferr, ErrStaleLoadJob) { + xlog.Debug("Failed to record the load placement", "model", ref.TrackingKey, "error", ferr) + } + cancelFlush() + } - stopHeartbeat() + cause := context.Cause(ownerCtx) + lost := errors.Is(cause, ErrStaleLoadJob) + leaseExpired := errors.Is(cause, ErrLoadLeaseExpired) + opLost := errors.Is(cause, ErrLoadOperationLost) + if opLost { + err = fmt.Errorf("loading model %s: %w", ref.TrackingKey, ErrLoadOperationLost) + } + if leaseExpired { + err = fmt.Errorf("loading model %s: %w", ref.TrackingKey, ErrLoadLeaseExpired) + } + // Bookkeeping must survive the owner context, which may be exactly what + // just ended. + bookCtx, cancelBook := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second) + defer cancelBook() - // Bookkeeping must survive the load context, which may be exactly what - // just expired. - bookCtx, cancelBook := context.WithTimeout(context.WithoutCancel(parent), 30*time.Second) - defer cancelBook() - - if err != nil { - xlog.Error("Cold load job failed", "model", trackingKey, "error", err) - if ferr := r.registry.FailLoadJob(bookCtx, trackingKey, err.Error()); ferr != nil { - xlog.Warn("Failed to record cold load failure", "model", trackingKey, "error", ferr) + switch { + case lost: + xlog.Warn("Cold load stopped: its job now belongs to another attempt", "model", ref.TrackingKey) + r.closeLoadWaiters(loadWaiterKey(ref)) + // The job was replaced or cancelled, but this attempt's remote work may + // still run. Stop it by its operation id. The confirmation shortens the + // stop window of a cancelled job. + r.stopLoadWork(bookCtx, ref, nodeID, replica, phase.backendAddress(), legacy) + return fmt.Errorf("loading model %s: %w", ref.TrackingKey, ErrStaleLoadJob) + case err != nil: + xlog.Error("Cold load job failed", "model", ref.TrackingKey, "error", err) + // Work may outlive the failure when we gave up on it (deadline, cancel, + // lost lease) once a node was chosen. An answer from the backend, or a + // failure before any node was chosen, means nothing is left running. + mayRun := phase.snapshot().State != LoadJobStatePending && (leaseExpired || loadAbandonedOnWorker(err)) + if ferr := r.registry.FailLoadJob(bookCtx, ref, err.Error(), mayRun); ferr != nil { + if errors.Is(ferr, ErrStaleLoadJob) { + r.closeLoadWaiters(loadWaiterKey(ref)) + return fmt.Errorf("loading model %s: %w (load error: %v)", ref.TrackingKey, ErrStaleLoadJob, err) } - r.closeLoadWaiters(trackingKey) - // Keep the row briefly so a request arriving right now reports this - // failure instead of starting a duplicate load. Deleting it - // immediately turns a failure into a retry storm. - time.AfterFunc(loadJobFailureGrace, func() { - delCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - if derr := r.registry.DeleteLoadJob(delCtx, trackingKey); derr != nil { - xlog.Warn("Failed to clear failed cold load job", "model", trackingKey, "error", derr) - } - }) + xlog.Warn("Failed to record cold load failure", "model", ref.TrackingKey, "error", ferr) + } + r.closeLoadWaiters(loadWaiterKey(ref)) + switch { + case !mayRun: + // The backend answered, so the work ended. Stop watching it. + r.completeLoadOperation(bookCtx, ref, nodeID) + default: + r.stopLoadWork(bookCtx, ref, nodeID, replica, phase.backendAddress(), legacy) + } + // The row stays until its stop deadline so a request arriving right + // now reports this failure instead of starting a duplicate load. The + // timer only tidies the table: a restart loses it, and the next claim + // replaces a failed row past its deadline without it. + time.AfterFunc(r.failedJobTidyDelay(mayRun), func() { + delCtx, cancelDel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancelDel() + if derr := r.registry.DeleteFailedLoadJob(delCtx, ref); derr != nil && !errors.Is(derr, ErrStaleLoadJob) { + xlog.Warn("Failed to clear failed cold load job", "model", ref.TrackingKey, "error", derr) + } + }) + return err + } + // The load succeeded: end the operation on the worker before the job goes, + // so the watchdog stops watching a backend that now serves. + r.completeLoadOperation(bookCtx, ref, nodeID) + return r.finishLoadJob(bookCtx, ref) +} + +// completeLoadOperation tells the worker the load finished. It retries a few +// times, and a loss is not fatal: the worker leaves a backend that already +// answers READY running when an operation expires. +func (r *SmartRouter) completeLoadOperation(ctx context.Context, ref LoadJobRef, nodeID string) { + if r.unloader == nil || nodeID == "" { + return + } + var lastErr error + for range 3 { + if _, lastErr = r.unloader.OperationControl(nodeID, workerctl.OperationRequest{Complete: []string{ref.Generation}}); lastErr == nil { return } + select { + case <-ctx.Done(): + return + case <-time.After(500 * time.Millisecond): + } + } + xlog.Warn("Could not complete the load operation on the worker", "node", nodeID, "model", ref.TrackingKey, "error", lastErr) +} - r.finishLoadJob(bookCtx, trackingKey) - }() +// stopLoadWork stops the remote work of one attempt through the single stop +// path, and records the outcome on the job. Without a node there is nothing to +// stop yet; the worker's own watchdog bounds any install already in flight. +func (r *SmartRouter) stopLoadWork(ctx context.Context, ref LoadJobRef, nodeID string, replica int, addr string, legacy bool) { + if r.unloader == nil || nodeID == "" { + return + } + reg, ok := r.registry.(loadAttemptRegistry) + if !ok { + return + } + StopLoadAttempt(ctx, reg, r.unloader, ref, nodeID, replica, addr, legacy) } // finishLoadJob ends a job that succeeded. The NodeModel row (state `loaded`) // is the record from here, so the job row is dropped BEFORE waiters are woken: -// they re-run the warm path and must not find a job that is really done. -func (r *SmartRouter) finishLoadJob(ctx context.Context, trackingKey string) { - if err := r.registry.DeleteLoadJob(ctx, trackingKey); err != nil { - xlog.Warn("Failed to clear completed cold load job", "model", trackingKey, "error", err) +// they re-run the warm path and must not find a job that is really done. A +// stale ref deletes nothing and wakes only its own generation's waiters. +func (r *SmartRouter) finishLoadJob(ctx context.Context, ref LoadJobRef) error { + err := r.registry.DeleteLoadJob(ctx, ref) + if err != nil { + xlog.Warn("Failed to clear completed cold load job", "model", ref.TrackingKey, "error", err) } - r.closeLoadWaiters(trackingKey) + r.closeLoadWaiters(loadWaiterKey(ref)) + if errors.Is(err, ErrStaleLoadJob) { + return fmt.Errorf("loading model %s: %w", ref.TrackingKey, err) + } + return nil } // startLoadJobHeartbeat keeps the job row's liveness and progress fresh while @@ -192,7 +300,15 @@ func (r *SmartRouter) finishLoadJob(ctx context.Context, trackingKey string) { // row when bytes moved would look orphaned and be reclaimed mid-load. Byte // progress is copied in from the staging tracker, which already debounces the // per-chunk callbacks, so the row is written at most once per interval. -func (r *SmartRouter) startLoadJobHeartbeat(parent context.Context, trackingKey string, phase *loadPhaseReporter) func() { +func (r *SmartRouter) startLoadJobHeartbeat(parent context.Context, ref LoadJobRef, phase *loadPhaseReporter, abort context.CancelCauseFunc) func() { + var renewing atomic.Bool + var unknown atomic.Int32 + ticks := 0 + every := r.opRenewEvery + if every <= 0 { + every = loadOpRenewEvery + } + trackingKey := ref.TrackingKey done := make(chan struct{}) stopped := make(chan struct{}) @@ -201,12 +317,21 @@ func (r *SmartRouter) startLoadJobHeartbeat(parent context.Context, trackingKey ticker := time.NewTicker(loadJobHeartbeatInterval) defer ticker.Stop() var startedAt time.Time + // lastRenewed is monotonic, so the owner's own deadline does not move + // with the wall clock. + lastRenewed := time.Now() for { select { case <-done: return + case <-parent.Done(): + return case <-ticker.C: u := phase.snapshot() + ticks++ + if ticks%every == 1 || every == 1 { + r.renewLoadOperation(parent, ref, phase, &renewing, &unknown, abort) + } if st := r.stagingTracker.Get(trackingKey); st != nil { u.BytesSent, u.TotalBytes = st.BytesSent, st.TotalBytes u.FileIndex, u.TotalFiles = st.FileIndex, st.TotalFiles @@ -216,10 +341,27 @@ func (r *SmartRouter) startLoadJobHeartbeat(parent context.Context, trackingKey u.StartedAt = startedAt } ctx, cancel := context.WithTimeout(context.WithoutCancel(parent), loadJobHeartbeatInterval*5) - if err := r.registry.UpdateLoadJob(ctx, trackingKey, u); err != nil { - xlog.Debug("Failed to heartbeat cold load job", "model", trackingKey, "error", err) - } + err := r.registry.UpdateLoadJob(ctx, ref, u) cancel() + if errors.Is(err, ErrStaleLoadJob) { + // Zero rows: the job is gone or another generation holds + // it. Stop the work instead of finishing a load nobody + // owns. + abort(ErrStaleLoadJob) + return + } + if err == nil { + lastRenewed = time.Now() + continue + } + xlog.Debug("Failed to heartbeat cold load job", "model", trackingKey, "error", err) + // A failed renewal is not fatal by itself. The owner keeps + // working until the lease it last extended has run out, then + // stops: it must not outlive a lease it cannot extend. + if time.Since(lastRenewed) >= r.loadLeaseTTL() { + abort(ErrLoadLeaseExpired) + return + } } } }() @@ -230,70 +372,173 @@ func (r *SmartRouter) startLoadJobHeartbeat(parent context.Context, trackingKey } } -// waitForLoadJob blocks until the cold load of trackingKey reaches a terminal -// state, the job's failure is known, or the caller gives up. +// waitForLoadJob blocks until the cold load of ref reaches a terminal state, +// the job's failure is known, or the caller gives up. // // Waiters share one broadcast rather than an ordered queue: they all want the // identical outcome — the model loaded — so ordering them would add fairness -// machinery that changes no result. The local channel wakes same-replica -// waiters instantly; the DB poll is the authority, because a waiter on another -// replica has no channel to close and NATS broadcasts are fire-and-forget, so a -// missed terminal event must not strand it. -func (r *SmartRouter) waitForLoadJob(ctx context.Context, trackingKey string, waiter <-chan struct{}) error { +// machinery that changes no result. The local channel is only a hint that wakes +// same-replica waiters early; the DB is the authority, because a waiter on +// another replica has no channel to close and broadcasts are +// fire-and-forget, so a missed terminal event must not strand it. +// +// The waiter is bound to one generation. When the job row is gone, or belongs to +// another generation, the attempt it waited for is over: the caller re-runs the +// warm path and, if the model is still missing, claims again. +func (r *SmartRouter) waitForLoadJob(ctx context.Context, ref LoadJobRef, waiter <-chan struct{}) error { + defer r.releaseLoadWaiter(loadWaiterKey(ref), waiter) + + // The registration may have come after the job ended; ask the authority + // before sleeping on the hint. + if done, err := r.checkLoadWait(ctx, ref); done { + return err + } + ticker := time.NewTicker(loadJobPollInterval) defer ticker.Stop() for { select { case <-waiter: - return nil + waiter = nil // closed: do not spin on it + if done, err := r.checkLoadWait(ctx, ref); done { + return err + } case <-ctx.Done(): // The client gave up. The job is unaffected: it is owned by the job // record, not by this request. return ctx.Err() case <-ticker.C: - job, err := r.registry.GetLoadJob(ctx, trackingKey) - if err != nil { - xlog.Debug("Polling the model load job failed", "model", trackingKey, "error", err) - continue - } - if job == nil { - // Terminal: either it succeeded, or it was reaped. Either way - // the caller re-checks the warm path. - return nil - } - if job.State == LoadJobStateFailed { - return fmt.Errorf("loading model %s: %s", trackingKey, job.LastError) + if done, err := r.checkLoadWait(ctx, ref); done { + return err } } } } -// loadWaiterChan returns the broadcast channel for trackingKey, creating it on -// first use. Same shape as advisorylock.localLocks: N local requests share one -// wait and wake together. -func (r *SmartRouter) loadWaiterChan(trackingKey string) <-chan struct{} { +// checkLoadWait reads the job row and reports whether the wait is over. +func (r *SmartRouter) checkLoadWait(ctx context.Context, ref LoadJobRef) (bool, error) { + job, err := r.registry.GetLoadJob(ctx, ref.TrackingKey) + if err != nil { + xlog.Debug("Polling the model load job failed", "model", ref.TrackingKey, "error", err) + return false, nil + } + if job == nil || job.Generation != ref.Generation { + // Terminal: it succeeded, was reaped, or was replaced. Either way the + // caller re-checks the warm path. + return true, nil + } + if job.State == LoadJobStateFailed { + return true, NewLoadHeldError(job) + } + return false, nil +} + +// loadWaiter is the shared wake-up channel for one generation, with a count of +// the requests registered on it. +type loadWaiter struct { + ch chan struct{} + refs int +} + +// loadWaiterKey keys waiters by generation, not by model, so a late finish of +// one attempt cannot wake the waiters of the next. +func loadWaiterKey(ref LoadJobRef) string { return ref.TrackingKey + "\x00" + ref.Generation } + +// loadWaiterChan registers a waiter and returns the broadcast channel for key, +// creating it on first use. Same shape as advisorylock.localLocks: N local +// requests share one wait and wake together. +func (r *SmartRouter) loadWaiterChan(key string) <-chan struct{} { r.loadWaitersMu.Lock() defer r.loadWaitersMu.Unlock() if r.loadWaiters == nil { - r.loadWaiters = map[string]chan struct{}{} + r.loadWaiters = map[string]*loadWaiter{} } - ch, ok := r.loadWaiters[trackingKey] + w, ok := r.loadWaiters[key] if !ok { - ch = make(chan struct{}) - r.loadWaiters[trackingKey] = ch + w = &loadWaiter{ch: make(chan struct{})} + r.loadWaiters[key] = w } - return ch + w.refs++ + return w.ch } -// closeLoadWaiters wakes every local waiter on trackingKey. A waiter that -// registers after this sees a fresh channel and falls back to the DB poll. -func (r *SmartRouter) closeLoadWaiters(trackingKey string) { +// releaseLoadWaiter drops one registration. Without it a waiter that gives up +// before the job ends would leave its entry in the map for ever. The channel +// identity check keeps a late release from deleting a newer registration. +func (r *SmartRouter) releaseLoadWaiter(key string, ch <-chan struct{}) { r.loadWaitersMu.Lock() - ch, ok := r.loadWaiters[trackingKey] - delete(r.loadWaiters, trackingKey) - r.loadWaitersMu.Unlock() - if ok { - close(ch) + defer r.loadWaitersMu.Unlock() + if w := r.loadWaiters[key]; w != nil && w.ch == ch { + w.refs-- + if w.refs <= 0 { + delete(r.loadWaiters, key) + } } } + +// closeLoadWaiters wakes every local waiter on key. A waiter that registers +// after this sees a fresh channel and falls back to the DB check. +func (r *SmartRouter) closeLoadWaiters(key string) { + r.loadWaitersMu.Lock() + w, ok := r.loadWaiters[key] + delete(r.loadWaiters, key) + r.loadWaitersMu.Unlock() + if ok { + close(w.ch) + } +} + +// loadOpRenewEvery is how many heartbeat ticks pass between operation renewals +// on the worker. Renewing every few seconds is far inside the worker's kill TTL +// and spares the bus a request per second per load. +const loadOpRenewEvery = 5 + +// loadOpLostAfter is how many renewals in a row the worker must answer with +// "unknown" before the owner gives up on the operation. +const loadOpLostAfter = 3 + +// renewLoadOperation extends the worker's lease on the load. It runs off the +// heartbeat goroutine, so a slow worker cannot delay the database lease, and at +// most one is in flight. A failure costs nothing until the worker's kill TTL. +func (r *SmartRouter) renewLoadOperation(ctx context.Context, ref LoadJobRef, phase *loadPhaseReporter, busy *atomic.Bool, unknown *atomic.Int32, abort context.CancelCauseFunc) { + nodeID, _, _ := phase.placement() + if r.unloader == nil || nodeID == "" || !busy.CompareAndSwap(false, true) { + return + } + go func() { + defer busy.Store(false) + reply, err := r.unloader.OperationControl(nodeID, workerctl.OperationRequest{Renew: []string{ref.Generation}}) + switch { + case err != nil: + xlog.Debug("Failed to renew the load operation", "node", nodeID, "model", ref.TrackingKey, "error", err) + case len(reply.Unknown) > 0: + // The worker answered and does not know the operation: it + // restarted, or its watchdog already ended it. Either way the work + // is gone, and nothing will ever finish this load. One miss can be an + // install still in flight, so it takes several in a row. + if n := unknown.Add(1); n >= loadOpLostAfter { + xlog.Warn("The worker no longer knows this load operation; failing the load", "node", nodeID, "model", ref.TrackingKey) + abort(ErrLoadOperationLost) + } + default: + unknown.Store(0) + } + }() +} + +func (r *SmartRouter) loadLeaseTTL() time.Duration { + if r.leaseTTL > 0 { + return r.leaseTTL + } + return loadJobLeaseTTL +} + +// failedJobTidyDelay is when the in-process timer tries to remove a failed row. +// It only has to be after the stop deadline the row was given. +func (r *SmartRouter) failedJobTidyDelay(mayRun bool) time.Duration { + if mayRun { + return loadJobStopWindow + time.Second + } + return loadJobFailureReport + time.Second +} diff --git a/core/services/nodes/load_job_sqlite_test.go b/core/services/nodes/load_job_sqlite_test.go new file mode 100644 index 000000000..3c71827a0 --- /dev/null +++ b/core/services/nodes/load_job_sqlite_test.go @@ -0,0 +1,139 @@ +// SPDX-License-Identifier: MIT +package nodes + +import ( + "context" + "path/filepath" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +// legacyLoadJob is the job table as a binary without the generation column made +// it. +type legacyLoadJob struct { + TrackingKey string `gorm:"primaryKey;size:255"` + State string `gorm:"size:16;not null;index"` + OwnerReplica string `gorm:"size:64"` + NodeID string `gorm:"size:36"` + NodeName string `gorm:"size:255"` + ReplicaIndex int + BytesSent int64 + TotalBytes int64 + FileIndex int + TotalFiles int + LastError string `gorm:"type:text"` + StartedAt time.Time + CreatedAt time.Time + UpdatedAt time.Time + LastProgress time.Time `gorm:"index"` +} + +func (legacyLoadJob) TableName() string { return "model_load_jobs" } + +// A single-process deployment keeps its registry in SQLite. The migration and +// the job rules must hold there too, not only on PostgreSQL. These specs need +// no container. +var _ = Describe("Load jobs on SQLite", func() { + var ( + db *gorm.DB + ctx context.Context + ) + + BeforeEach(func() { + ctx = context.Background() + var err error + db, err = gorm.Open(sqlite.Open(filepath.Join(GinkgoT().TempDir(), "nodes.db")), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + }) + + It("gives legacy rows a generation of their own and fences them", func() { + Expect(db.Migrator().CreateTable(&legacyLoadJob{})).To(Succeed()) + for _, key := range []string{"legacy-a", "legacy-b"} { + Expect(db.Exec(`INSERT INTO model_load_jobs (tracking_key, state, owner_replica, last_progress, created_at, updated_at) + VALUES (?, 'staging', 'old-frontend', ?, ?, ?)`, key, time.Now(), time.Now(), time.Now()).Error).To(Succeed()) + } + + registry, err := NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + + var jobs []ModelLoadJob + Expect(db.Find(&jobs).Error).To(Succeed()) + Expect(jobs).To(HaveLen(2)) + Expect(jobs[0].Generation).ToNot(BeEmpty()) + Expect(jobs[1].Generation).ToNot(BeEmpty()) + Expect(jobs[0].Generation).ToNot(Equal(jobs[1].Generation)) + + // Running it again changes nothing. + before := jobs[0].Generation + _, err = NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + again, err := registry.GetLoadJob(ctx, jobs[0].TrackingKey) + Expect(err).ToNot(HaveOccurred()) + Expect(again.Generation).To(Equal(before)) + }) + + It("claims, heartbeats and releases a job, and fences a replaced attempt", func() { + registry, err := NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + + a, claimed, err := registry.ClaimLoadJob(ctx, "sqlite-model", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(registry.UpdateLoadJob(ctx, a.Ref(), LoadJobUpdate{State: LoadJobStateLoading})).To(Succeed()) + Expect(registry.DeleteLoadJob(ctx, a.Ref())).To(Succeed()) + + b, claimed, err := registry.ClaimLoadJob(ctx, "sqlite-model", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(registry.UpdateLoadJob(ctx, a.Ref(), LoadJobUpdate{})).To(MatchError(ErrStaleLoadJob)) + Expect(registry.FailLoadJob(ctx, b.Ref(), "boom", false)).To(Succeed()) + Expect(registry.DeleteLoadJob(ctx, b.Ref())).To(MatchError(ErrStaleLoadJob)) + }) + + // The lease rules read a clock that moves, so these specs move the stored + // deadlines instead of waiting. + It("applies the lease and stop window rules", func() { + registry, err := NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + past := time.Now().Add(-time.Second) + + dead, claimed, err := registry.ClaimLoadJob(ctx, "sqlite-lease", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + other, claimed, err := registry.ClaimLoadJob(ctx, "sqlite-lease", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse(), "a live lease holds the model") + Expect(other.State).ToNot(Equal(LoadJobStateFailed)) + + // The owner stops renewing: the claim fails the job and holds the model. + Expect(db.Exec("UPDATE model_load_jobs SET lease_until = ?", past).Error).To(Succeed()) + held, claimed, err := registry.ClaimLoadJob(ctx, "sqlite-lease", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + Expect(held.State).To(Equal(LoadJobStateFailed)) + Expect(held.StopDeadline).ToNot(BeNil()) + Expect(held.StopDeadline.After(time.Now().Add(loadJobStopWindow - 10*time.Second))).To(BeTrue()) + + // The stop window ends: the sweep releases it, and the next claim wins. + Expect(db.Exec("UPDATE model_load_jobs SET stop_deadline = ?", past).Error).To(Succeed()) + sweep, err := registry.SweepLoadJobs(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(sweep.Released).To(HaveLen(1)) + next, claimed, err := registry.ClaimLoadJob(ctx, "sqlite-lease", "frontend-b") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeTrue()) + Expect(next.Generation).ToNot(Equal(dead.Generation)) + + // A heartbeat renews the lease. + Expect(db.Exec("UPDATE model_load_jobs SET lease_until = ?", past).Error).To(Succeed()) + Expect(registry.UpdateLoadJob(ctx, next.Ref(), LoadJobUpdate{})).To(Succeed()) + live, err := registry.GetLoadJob(ctx, "sqlite-lease") + Expect(err).ToNot(HaveOccurred()) + Expect(live.LeaseUntil).ToNot(BeNil()) + Expect(live.LeaseUntil.After(time.Now())).To(BeTrue()) + }) +}) diff --git a/core/services/nodes/load_operation_control_conformance_test.go b/core/services/nodes/load_operation_control_conformance_test.go new file mode 100644 index 000000000..6ac3f0dd7 --- /dev/null +++ b/core/services/nodes/load_operation_control_conformance_test.go @@ -0,0 +1,254 @@ +// SPDX-License-Identifier: MIT +package nodes + +import ( + "context" + "encoding/json" + "errors" + "time" + + "github.com/nats-io/nats.go" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// sentRequest is one control request a carrier put on the wire, decoded. +type sentRequest struct { + Verb string + Timeout time.Duration + Install workerctl.BackendInstallRequest + Stop workerctl.ModelStopRequest + Op workerctl.OperationRequest + Unload workerctl.ModelUnloadRequest +} + +// loadOperationHarness drives one carrier for the conformance of +// LoadOperationControl. A second carrier plugs in by implementing it and being +// added to loadOperationCarriers. The harness knows how a failure looks on its +// own wire; the specs only say which of the four conditions it stands for. +type loadOperationHarness interface { + // Control is the carrier under test. + Control() LoadOperationControl + // NoRoute makes every later call find no route to the node. + NoRoute() + // TimesOut makes every later call get no reply in time. + TimesOut() + // WorkerRefuses makes the worker answer every call with its own refusal. + WorkerRefuses() + // WorkerAnswers makes the worker answer every call with success. + WorkerAnswers() + // Sent returns the requests the carrier sent, in order. + Sent() []sentRequest +} + +var loadOperationCarriers = map[string]func() loadOperationHarness{ + "the NATS carrier": newNATSLoadOperationHarness, +} + +type natsLoadOperationHarness struct { + mc *scriptedMessagingClient + adapter *RemoteUnloaderAdapter +} + +func newNATSLoadOperationHarness() loadOperationHarness { + mc := newScriptedMessagingClient() + return &natsLoadOperationHarness{mc: mc, adapter: NewRemoteUnloaderAdapter(&fakeModelLocator{}, mc, 3*time.Minute, 15*time.Minute)} +} + +func (h *natsLoadOperationHarness) Control() LoadOperationControl { return h.adapter } + +func (h *natsLoadOperationHarness) subjects() []string { + const node = conformanceNode + return []string{ + messaging.SubjectNodeBackendInstall(node), messaging.SubjectNodeModelStop(node), + messaging.SubjectNodeModelOp(node), messaging.SubjectNodeModelUnload(node), + } +} + +func (h *natsLoadOperationHarness) NoRoute() { + for _, s := range h.subjects() { + h.mc.scriptNoResponders(s) + } +} + +func (h *natsLoadOperationHarness) TimesOut() { + for _, s := range h.subjects() { + h.mc.scriptErr(s, nats.ErrTimeout) + } +} + +func (h *natsLoadOperationHarness) WorkerRefuses() { + const node = conformanceNode + h.mc.scriptReply(messaging.SubjectNodeBackendInstall(node), workerctl.BackendInstallReply{Success: false, Error: "disk full"}) + h.mc.scriptReply(messaging.SubjectNodeModelStop(node), workerctl.ModelStopReply{Matched: true, Error: "does not belong to operation"}) + h.mc.scriptReply(messaging.SubjectNodeModelOp(node), workerctl.OperationReply{Unknown: []string{"op"}}) + h.mc.scriptReply(messaging.SubjectNodeModelUnload(node), workerctl.ModelUnloadReply{Success: false, Error: "process was replaced during unload"}) +} + +func (h *natsLoadOperationHarness) WorkerAnswers() { + const node = conformanceNode + h.mc.scriptReply(messaging.SubjectNodeBackendInstall(node), workerctl.BackendInstallReply{Success: true, Address: "127.0.0.1:9001", ProcessInstance: "i", ReportsOperations: true}) + h.mc.scriptReply(messaging.SubjectNodeModelStop(node), workerctl.ModelStopReply{Matched: true, Terminated: true}) + h.mc.scriptReply(messaging.SubjectNodeModelOp(node), workerctl.OperationReply{Renewed: []string{"op"}, Completed: []string{"op"}}) + h.mc.scriptReply(messaging.SubjectNodeModelUnload(node), workerctl.ModelUnloadReply{Success: true}) +} + +func (h *natsLoadOperationHarness) Sent() []sentRequest { + h.mc.mu.Lock() + defer h.mc.mu.Unlock() + const node = conformanceNode + var out []sentRequest + for _, c := range h.mc.calls { + r := sentRequest{Timeout: c.Timeout} + switch c.Subject { + case messaging.SubjectNodeBackendInstall(node): + r.Verb = "install" + Expect(json.Unmarshal(c.Data, &r.Install)).To(Succeed()) + case messaging.SubjectNodeModelStop(node): + r.Verb = "stop" + Expect(json.Unmarshal(c.Data, &r.Stop)).To(Succeed()) + case messaging.SubjectNodeModelOp(node): + r.Verb = "op" + Expect(json.Unmarshal(c.Data, &r.Op)).To(Succeed()) + case messaging.SubjectNodeModelUnload(node): + r.Verb = "unload" + Expect(json.Unmarshal(c.Data, &r.Unload)).To(Succeed()) + } + out = append(out, r) + } + return out +} + +const conformanceNode = "11111111-2222-3333-4444-555555555555" + +// Every carrier of LoadOperationControl must pass these. They pin the contract +// in interfaces.go, so a carrier can be written from the interface alone. +var _ = Describe("LoadOperationControl conformance", func() { + for name, newHarness := range loadOperationCarriers { + Describe(name, func() { + var h loadOperationHarness + + BeforeEach(func() { h = newHarness() }) + + // call is one method, with the answer it should give a worker that + // answers. + type call struct { + verb string + run func() error + } + calls := func() []call { + c := h.Control() + replica := NodeModel{ModelName: "m", ReplicaIndex: 1, Address: "127.0.0.1:9001"} + return []call{ + {"install", func() error { + _, err := c.InstallBackendOp(conformanceNode, "llama-cpp", "m", "", 1, "", "op", time.Hour, nil) + return err + }}, + {"stop", func() error { + _, err := c.StopLoadOperation(context.Background(), conformanceNode, workerctl.ModelStopRequest{ + ModelName: "m", ProcessKey: "m#1", ExpectedAddress: "127.0.0.1:9001", OperationID: "op", Force: true}) + return err + }}, + {"op", func() error { + _, err := c.OperationControl(conformanceNode, workerctl.OperationRequest{Renew: []string{"op"}}) + return err + }}, + {"unload", func() error { return c.UnloadReplica(conformanceNode, replica) }}, + } + } + + It("reports no route as ErrNoRoute, and only that", func() { + h.NoRoute() + for _, c := range calls() { + err := c.run() + Expect(errors.Is(err, ErrNoRoute)).To(BeTrue(), "%s: %v", c.verb, err) + } + }) + + It("does not report a timeout as ErrNoRoute", func() { + h.TimesOut() + for _, c := range calls() { + err := c.run() + Expect(err).To(HaveOccurred(), c.verb) + Expect(errors.Is(err, ErrNoRoute)).To(BeFalse(), "%s: a slow worker is not an absent one", c.verb) + } + }) + + It("never reports a worker's own refusal as ErrNoRoute", func() { + h.WorkerRefuses() + for _, c := range calls() { + err := c.run() + Expect(errors.Is(err, ErrNoRoute)).To(BeFalse(), "%s: the worker answered, so it is present", c.verb) + } + }) + + It("hands a worker's refusal back in the reply, with no error", func() { + h.WorkerRefuses() + c := h.Control() + install, err := c.InstallBackendOp(conformanceNode, "b", "m", "", 0, "", "op", time.Hour, nil) + Expect(err).ToNot(HaveOccurred()) + Expect(install.Success).To(BeFalse()) + stop, err := c.StopLoadOperation(context.Background(), conformanceNode, workerctl.ModelStopRequest{ProcessKey: "m#0", OperationID: "op"}) + Expect(err).ToNot(HaveOccurred()) + Expect(stop.Error).ToNot(BeEmpty()) + ops, err := c.OperationControl(conformanceNode, workerctl.OperationRequest{Renew: []string{"op"}}) + Expect(err).ToNot(HaveOccurred()) + Expect(ops.Unknown).To(ConsistOf("op")) + }) + + It("gives every call its own bounded timeout, and renewals the shortest", func() { + h.WorkerAnswers() + for _, c := range calls() { + Expect(c.run()).To(Succeed(), c.verb) + } + for _, s := range h.Sent() { + Expect(s.Timeout).To(BeNumerically(">", 0), s.Verb) + if s.Verb == "op" { + Expect(s.Timeout).To(BeNumerically("<=", 5*time.Second), + "a renewal that takes longer than its cadence cannot hold the kill TTL") + } + } + }) + + It("carries the operation, the process and the deadline to the worker", func() { + h.WorkerAnswers() + for _, c := range calls() { + Expect(c.run()).To(Succeed(), c.verb) + } + sent := map[string]sentRequest{} + for _, s := range h.Sent() { + sent[s.Verb] = s + } + Expect(sent["install"].Install.OperationID).To(Equal("op")) + Expect(sent["install"].Install.DeadlineMs).To(Equal(time.Hour.Milliseconds()), "a duration, not a timestamp") + Expect(sent["stop"].Stop.OperationID).To(Equal("op")) + Expect(sent["stop"].Stop.ProcessKey).To(Equal("m#1")) + Expect(sent["stop"].Stop.ExpectedAddress).To(Equal("127.0.0.1:9001")) + Expect(sent["op"].Op.Renew).To(ConsistOf("op")) + Expect(sent["unload"].Unload.Address).To(Equal("127.0.0.1:9001")) + }) + + It("stops idempotently: a repeated stop is answered, not refused", func() { + h.WorkerAnswers() + c := h.Control() + for range 2 { + reply, err := c.StopLoadOperation(context.Background(), conformanceNode, workerctl.ModelStopRequest{ProcessKey: "m#0", OperationID: "op"}) + Expect(err).ToNot(HaveOccurred()) + Expect(reply.Terminated).To(BeTrue()) + } + }) + + It("refuses a stop with no operation id, and an unload with no address, without sending anything", func() { + h.WorkerAnswers() + c := h.Control() + _, err := c.StopLoadOperation(context.Background(), conformanceNode, workerctl.ModelStopRequest{ProcessKey: "m#0"}) + Expect(err).To(HaveOccurred()) + Expect(c.UnloadReplica(conformanceNode, NodeModel{ModelName: "m"})).To(Succeed()) + Expect(h.Sent()).To(BeEmpty(), "a carrier must never ask a worker to pick a process") + }) + }) + } +}) diff --git a/core/services/nodes/load_worker_ops_test.go b/core/services/nodes/load_worker_ops_test.go new file mode 100644 index 000000000..2f0a3ca16 --- /dev/null +++ b/core/services/nodes/load_worker_ops_test.go @@ -0,0 +1,448 @@ +// SPDX-License-Identifier: MIT +package nodes + +import ( + "context" + "errors" + "runtime" + "sync" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gorm.io/gorm" + + "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// fakeOpWorker is a worker as the controller sees it over the bus. It starts +// load operations, renews and completes them, and answers or ignores stops. It +// records every request, so a spec can assert on what was sent and on what was +// never sent. +type fakeOpWorker struct { + *fakeUnloader + + mu sync.Mutex + legacy bool // replies like a worker that predates operations + stopMode string // "ack", "hang" or "refuse" + installs []workerctl.BackendInstallRequest + renews []string + complete []string + stops []workerctl.ModelStopRequest + // exact records the address-addressed stops sent to a legacy worker, and + // exactFails makes them fail like a worker that never heard of the verb. + exact []NodeModel + exactFails bool + // forgetOps makes the worker answer every renewal with "unknown", as a + // worker that restarted and lost its operations does. + forgetOps bool +} + +func (w *fakeOpWorker) StopModelReplica(_ context.Context, _ string, replica NodeModel, _ bool) (workerctl.ModelStopReply, error) { + w.mu.Lock() + defer w.mu.Unlock() + w.exact = append(w.exact, replica) + if w.exactFails { + return workerctl.ModelStopReply{}, ErrNoRoute + } + return workerctl.ModelStopReply{Matched: true, Terminated: true}, nil +} + +func (w *fakeOpWorker) InstallBackendOp(nodeID, backend, modelID, galleries string, replica int, opID, operationID string, deadline time.Duration, progress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { + w.mu.Lock() + w.installs = append(w.installs, workerctl.BackendInstallRequest{ModelID: modelID, ReplicaIndex: int32(replica), OperationID: operationID, DeadlineMs: deadline.Milliseconds()}) + w.mu.Unlock() + reply, err := w.fakeUnloader.InstallBackend(nodeID, backend, modelID, galleries, "", "", "", replica, opID, progress) + if reply != nil { + copy := *reply + copy.ProcessInstance, copy.ReportsOperations = "", false + if !w.legacy { + copy.ProcessInstance, copy.ReportsOperations = "instance-1", true + } + reply = © + } + return reply, err +} + +func (w *fakeOpWorker) StopLoadOperation(_ context.Context, _ string, req workerctl.ModelStopRequest) (workerctl.ModelStopReply, error) { + w.mu.Lock() + w.stops = append(w.stops, req) + mode := w.stopMode + w.mu.Unlock() + switch mode { + case "hang": + return workerctl.ModelStopReply{}, errors.New("nats: timeout") + case "refuse": + return workerctl.ModelStopReply{Matched: true, Error: "does not belong to operation"}, nil + } + return workerctl.ModelStopReply{Matched: true, Terminated: true, ProcessKey: req.ProcessKey}, nil +} + +func (w *fakeOpWorker) OperationControl(_ string, req workerctl.OperationRequest) (*workerctl.OperationReply, error) { + w.mu.Lock() + defer w.mu.Unlock() + w.renews = append(w.renews, req.Renew...) + w.complete = append(w.complete, req.Complete...) + if w.forgetOps { + return &workerctl.OperationReply{Unknown: req.Renew, Completed: req.Complete}, nil + } + return &workerctl.OperationReply{Renewed: req.Renew, Completed: req.Complete}, nil +} + +func (w *fakeOpWorker) setStopMode(mode string) { + w.mu.Lock() + w.stopMode = mode + w.mu.Unlock() +} + +func (w *fakeOpWorker) snapshot() (installs []workerctl.BackendInstallRequest, renews, complete []string, stops []workerctl.ModelStopRequest) { + w.mu.Lock() + defer w.mu.Unlock() + return append(installs, w.installs...), append(renews, w.renews...), append(complete, w.complete...), append(stops, w.stops...) +} + +var _ = Describe("Load operations on the worker", func() { + var ( + db *gorm.DB + registry *NodeRegistry + ctx context.Context + worker *fakeOpWorker + backend *stubBackend + router *SmartRouter + node *BackendNode + ) + + BeforeEach(func() { + if runtime.GOOS == "darwin" { + Skip("testcontainers requires Docker, not available on macOS CI") + } + db = testutil.SetupTestDB() + var err error + registry, err = NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + ctx = context.Background() + node = &BackendNode{Name: "worker-1", NodeType: NodeTypeBackend, Address: "10.0.0.1:50051", + TotalVRAM: 64_000_000_000, AvailableVRAM: 64_000_000_000} + Expect(registry.Register(ctx, node, true)).To(Succeed()) + backend = &stubBackend{healthResult: true, loadResult: &pb.Result{Success: true}} + worker = &fakeOpWorker{fakeUnloader: &fakeUnloader{installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}}, stopMode: "ack"} + router = NewSmartRouter(registry, SmartRouterOptions{Unloader: worker, ClientFactory: &stubClientFactory{client: backend}, DB: db}) + }) + + opts := &pb.ModelOptions{Model: "models/m.gguf"} + route := func(model string) error { + res, err := router.Route(ctx, model, "models/m.gguf", "llama-cpp", "", opts, false) + if res != nil { + res.Release() + } + return err + } + jobOf := func(model string) *ModelLoadJob { + job, err := registry.GetLoadJob(ctx, model) + Expect(err).ToNot(HaveOccurred()) + return job + } + secondsUntil := func(model, column string) float64 { + var secs float64 + Expect(db.Raw("SELECT EXTRACT(EPOCH FROM ("+column+" - now())) FROM model_load_jobs WHERE tracking_key = ?", model).Scan(&secs).Error).To(Succeed()) + return secs + } + release := func(model string) { + Expect(db.Exec("UPDATE model_load_jobs SET stop_deadline = now() - interval '1 second' WHERE tracking_key = ?", model).Error).To(Succeed()) + } + + It("starts the load as a bounded operation, renews it, and completes it on success", func() { + hold := make(chan struct{}) + worker.installHook = func() { <-hold } + done := make(chan error, 1) + go func() { defer GinkgoRecover(); done <- route("ops-ok") }() + + var generation string + Eventually(func() string { + if job := jobOf("ops-ok"); job != nil { + generation = job.Generation + } + return generation + }, 5*time.Second, 50*time.Millisecond).ShouldNot(BeEmpty()) + // Held inside the install, the heartbeat renews on the worker. + Eventually(func() []string { _, renews, _, _ := worker.snapshot(); return renews }, 10*time.Second, 100*time.Millisecond). + Should(ContainElement(generation)) + close(hold) + Eventually(done, 15*time.Second).Should(Receive(BeNil())) + + installs, _, complete, stops := worker.snapshot() + Expect(installs).To(HaveLen(1)) + Expect(installs[0].OperationID).To(Equal(generation), "the operation id is the job generation") + Expect(installs[0].DeadlineMs).To(BeNumerically(">", 0)) + Expect(complete).To(ContainElement(generation), "a load that finished must stop being watched") + Expect(stops).To(BeEmpty()) + Eventually(func() *ModelLoadJob { return jobOf("ops-ok") }, 5*time.Second).Should(BeNil()) + }) + + It("stops the operation by id when a load times out, and shortens the hold once the worker acknowledges", func() { + backend.loadErr = context.DeadlineExceeded + Expect(route("ops-timeout")).To(HaveOccurred()) + + // The owner records the stop after the caller has its answer. + Eventually(func() bool { return jobOf("ops-timeout").OpConfirmed }, 5*time.Second, 50*time.Millisecond).Should(BeTrue()) + job := jobOf("ops-timeout") + Expect(job.State).To(Equal(LoadJobStateFailed)) + _, _, _, stops := worker.snapshot() + Expect(stops).To(HaveLen(1)) + Expect(stops[0].OperationID).To(Equal(job.Generation)) + Expect(stops[0].ProcessKey).To(Equal("ops-timeout#0")) + Expect(secondsUntil("ops-timeout", "stop_deadline")).To(BeNumerically("<=", loadJobFailureReport.Seconds()+3), + "an acknowledged stop frees the model in seconds, not after the full stop window") + + // A request inside the report window reads the cause as a 503 answer. + err := route("ops-timeout") + var held *ModelLoadingError + Expect(errors.As(err, &held)).To(BeTrue()) + Expect(held.RetryAfter).To(BeNumerically(">=", time.Second)) + Expect(held.Status.State).To(Equal(LoadJobStateFailed)) + + release("ops-timeout") + backend.mu.Lock() + backend.loadErr = nil + backend.mu.Unlock() + Expect(route("ops-timeout")).To(Succeed(), "the model loads again with no manual cleanup") + }) + + It("holds the model for the stop window while a worker stays silent, retries the stop, and releases at the deadline", func() { + worker.setStopMode("hang") + backend.loadErr = context.DeadlineExceeded + Expect(route("ops-hang")).To(HaveOccurred()) + + Eventually(func() int { _, _, _, stops := worker.snapshot(); return len(stops) }, 5*time.Second, 50*time.Millisecond).Should(Equal(1)) + job := jobOf("ops-hang") + Expect(job.OpConfirmed).To(BeFalse()) + Expect(job.NodeID).To(Equal(node.ID), "the stop needs the node, even when the load failed within one heartbeat") + Expect(secondsUntil("ops-hang", "stop_deadline")).To(BeNumerically("~", loadJobStopWindow.Seconds(), 3)) + var stops []workerctl.ModelStopRequest + + // Not before the deadline. + _, claimed, err := registry.ClaimLoadJob(ctx, "ops-hang", "other-frontend") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse()) + + // The reconciler retries the stop on each pass while the worker is silent. + rc := NewReplicaReconciler(ReplicaReconcilerOptions{Registry: registry, DB: db, Unloader: worker}) + rc.reclaimAbandonedLoads(ctx) + _, _, _, stops = worker.snapshot() + Expect(stops).To(HaveLen(2)) + Expect(jobOf("ops-hang").OpConfirmed).To(BeFalse()) + + // The worker answers at last. The hold shrinks to the report window. + worker.setStopMode("ack") + rc.reclaimAbandonedLoads(ctx) + Expect(jobOf("ops-hang").OpConfirmed).To(BeTrue()) + Expect(secondsUntil("ops-hang", "stop_deadline")).To(BeNumerically("<=", loadJobFailureReport.Seconds()+3)) + }) + + It("stops a legacy worker's process by exact address only, never by name", func() { + worker.legacy = true + backend.loadErr = context.DeadlineExceeded + Expect(route("ops-legacy")).To(HaveOccurred()) + + Eventually(func() bool { return jobOf("ops-legacy").OpConfirmed }, 5*time.Second, 50*time.Millisecond).Should(BeTrue()) + worker.mu.Lock() + exact := append([]NodeModel(nil), worker.exact...) + worker.mu.Unlock() + Expect(exact).To(HaveLen(1)) + Expect(exact[0].Address).To(Equal("10.0.0.1:9001")) + Expect(exact[0].ReplicaIndex).To(BeZero()) + _, _, _, stops := worker.snapshot() + Expect(stops).To(BeEmpty(), "a legacy worker is not sent operation stops") + + worker.fakeUnloader.mu.Lock() + defer worker.fakeUnloader.mu.Unlock() + Expect(worker.fakeUnloader.stopCalls).To(BeEmpty(), "no backend.stop by model name, ever") + Expect(worker.fakeUnloader.unloadCalls).To(BeEmpty()) + }) + + It("holds the model for the load deadline when a legacy worker cannot be stopped at all", func() { + worker.legacy = true + worker.exactFails = true + backend.loadErr = context.DeadlineExceeded + Expect(route("ops-legacy-silent")).To(HaveOccurred()) + + Eventually(func() float64 { return secondsUntil("ops-legacy-silent", "stop_deadline") }, 5*time.Second, 50*time.Millisecond). + Should(BeNumerically("~", loadJobLegacyStopWindow.Seconds(), 5)) + Expect(jobOf("ops-legacy-silent").OpConfirmed).To(BeFalse(), "a legacy worker cannot confirm a stop") + worker.fakeUnloader.mu.Lock() + defer worker.fakeUnloader.mu.Unlock() + Expect(worker.fakeUnloader.stopCalls).To(BeEmpty()) + }) + + It("lets a failure the backend answered end the operation instead of stopping the process", func() { + backend.loadResult = &pb.Result{Success: false, Message: "unsupported architecture"} + Expect(route("ops-answered")).To(HaveOccurred()) + + job := jobOf("ops-answered") + Expect(job.OpConfirmed).To(BeTrue()) + Eventually(func() []string { _, _, complete, _ := worker.snapshot(); return complete }, 5*time.Second, 50*time.Millisecond).Should(ContainElement(job.Generation)) + _, _, complete, stops := worker.snapshot() + Expect(stops).To(BeEmpty(), "the backend answered, so its process is idle and stays warm") + Expect(complete).To(ContainElement(job.Generation)) + }) + + It("confirms every failed attempt on a node when its worker restarts", func() { + Expect(registry.ObserveWorkerIncarnation(ctx, node.ID, "boot-1")).To(Succeed()) + job, _, err := registry.ClaimLoadJob(ctx, "ops-restart", "frontend-a") + Expect(err).ToNot(HaveOccurred()) + Expect(registry.UpdateLoadJob(ctx, job.Ref(), LoadJobUpdate{State: LoadJobStateLoading, NodeID: node.ID, NodeName: node.Name})).To(Succeed()) + Expect(registry.FailLoadJob(ctx, job.Ref(), "context deadline exceeded", true)).To(Succeed()) + Expect(jobOf("ops-restart").OpConfirmed).To(BeFalse()) + + // The same incarnation proves nothing. + Expect(registry.ObserveWorkerIncarnation(ctx, node.ID, "boot-1")).To(Succeed()) + Expect(jobOf("ops-restart").OpConfirmed).To(BeFalse()) + + // A new one proves the old process, and every operation in it, ended. + Expect(registry.ObserveWorkerIncarnation(ctx, node.ID, "boot-2")).To(Succeed()) + Expect(jobOf("ops-restart").OpConfirmed).To(BeTrue()) + Expect(secondsUntil("ops-restart", "stop_deadline")).To(BeNumerically("<=", loadJobFailureReport.Seconds()+3)) + }) + + It("stops the owner of a cancelled load, and the model loads again once the worker confirmed", func() { + hold := make(chan struct{}) + worker.installHook = func() { <-hold } + defer func() { + select { + case <-hold: + default: + close(hold) + } + }() + done := make(chan error, 1) + go func() { defer GinkgoRecover(); done <- route("ops-cancel") }() + + var job *ModelLoadJob + Eventually(func() string { + job = jobOf("ops-cancel") + if job == nil { + return "" + } + return job.NodeID + }, 10*time.Second, 50*time.Millisecond).ShouldNot(BeEmpty(), "the heartbeat records where the load runs") + + result, err := (&LoadCancelService{Registry: registry, Stopper: worker}).Cancel(ctx, job.Ref()) + Expect(err).ToNot(HaveOccurred()) + Expect(result.State).To(Equal(LoadCancelStopped)) + + // The owner notices at its next heartbeat and stops its own work. + var routeErr error + Eventually(done, 15*time.Second).Should(Receive(&routeErr)) + Expect(routeErr).To(HaveOccurred()) + failed := jobOf("ops-cancel") + Expect(failed.CancelRequested).To(BeTrue()) + Expect(failed.Generation).To(Equal(job.Generation)) + + release("ops-cancel") + close(hold) + Expect(route("ops-cancel")).To(Succeed()) + }) + + It("fails a load promptly when the worker no longer knows its operation", func() { + // A worker that was killed and restarted has no record of the load. Its + // predecessor's backend is gone, so waiting out the load budget buys + // nothing. + worker.forgetOps = true + router.opRenewEvery = 1 + hold := make(chan struct{}) + defer close(hold) + worker.installHook = func() { <-hold } + + started := time.Now() + err := route("ops-lost") + + Expect(err).To(HaveOccurred()) + Expect(time.Since(started)).To(BeNumerically("<", 30*time.Second), "not the five minute load budget") + Eventually(func() bool { j := jobOf("ops-lost"); return j != nil && j.OpConfirmed }, 5*time.Second, 50*time.Millisecond).Should(BeTrue(), + "the worker lost the operation, so its work ended") + job := jobOf("ops-lost") + Expect(job.State).To(Equal(LoadJobStateFailed)) + Expect(job.LastError).To(ContainSubstring("operation")) + Expect(secondsUntil("ops-lost", "stop_deadline")).To(BeNumerically("<=", loadJobFailureReport.Seconds()+3)) + _, _, _, stops := worker.snapshot() + Expect(stops).To(BeEmpty(), "there is nothing left on the worker to stop") + }) + + Describe("cancelling a load", func() { + // The owner can be stuck inside a call that never returns, so the cancel + // itself has to clear the attempt's replica rows, not wait for the owner's + // cleanup or for a reconciler pass. + blockLoad := func() (release func()) { + hold := make(chan struct{}) + backend.loadHook = func(*pb.ModelOptions) { <-hold } + var once sync.Once + return func() { once.Do(func() { close(hold) }) } + } + replicaRows := func(model string) int64 { + var n int64 + Expect(db.Model(&NodeModel{}).Where("model_name = ?", model).Count(&n).Error).To(Succeed()) + return n + } + startAndPlace := func(model string) *ModelLoadJob { + go func() { defer GinkgoRecover(); _ = route(model) }() + var job *ModelLoadJob + Eventually(func() bool { + job = jobOf(model) + return job != nil && job.NodeID != "" && replicaRows(model) == 1 + }, 10*time.Second, 50*time.Millisecond).Should(BeTrue()) + Eventually(func() string { + var nm NodeModel + _ = db.First(&nm, "model_name = ?", model).Error + return nm.Address + }, 10*time.Second, 50*time.Millisecond).ShouldNot(BeEmpty(), "the replica row reaches the backend address") + return job + } + + It("clears the attempt's replica row at once, so the model loads again after the report window", func() { + defer blockLoad()() + job := startAndPlace("cancel-rows") + + result, err := (&LoadCancelService{Registry: registry, Stopper: worker}).Cancel(ctx, job.Ref()) + Expect(err).ToNot(HaveOccurred()) + Expect(result.State).To(Equal(LoadCancelStopped)) + + Expect(replicaRows("cancel-rows")).To(BeZero(), "a stale loading row holds the slot past the hold") + Expect(secondsUntil("cancel-rows", "stop_deadline")).To(BeNumerically("<=", loadJobFailureReport.Seconds()+3)) + }) + + It("stops a legacy worker's backend by exact address, not by an operation id", func() { + worker.legacy = true + defer blockLoad()() + job := startAndPlace("cancel-legacy") + Eventually(func() bool { return jobOf("cancel-legacy").LegacyWorker }, 5*time.Second, 50*time.Millisecond).Should(BeTrue()) + + result, err := (&LoadCancelService{Registry: registry, Stopper: worker}).Cancel(ctx, job.Ref()) + + Expect(err).ToNot(HaveOccurred()) + Expect(result.State).To(Equal(LoadCancelStopped)) + worker.mu.Lock() + defer worker.mu.Unlock() + Expect(worker.stops).To(BeEmpty(), "an old worker does not know operations, and answers an empty address with a mismatch") + Expect(worker.exact).To(HaveLen(1)) + Expect(worker.exact[0].Address).To(Equal("10.0.0.1:9001")) + }) + + It("holds a legacy worker's model for the load deadline when the stop is not acknowledged", func() { + worker.legacy = true + worker.exactFails = true + defer blockLoad()() + job := startAndPlace("cancel-legacy-silent") + Eventually(func() bool { return jobOf("cancel-legacy-silent").LegacyWorker }, 5*time.Second, 50*time.Millisecond).Should(BeTrue()) + + result, err := (&LoadCancelService{Registry: registry, Stopper: worker}).Cancel(ctx, job.Ref()) + + Expect(err).ToNot(HaveOccurred()) + Expect(result.State).To(Equal(LoadCancelStopping)) + Expect(secondsUntil("cancel-legacy-silent", "stop_deadline")).To(BeNumerically("~", loadJobLegacyStopWindow.Seconds(), 5), + "the backend may still run, and nothing sooner bounds it") + Expect(jobOf("cancel-legacy-silent").OpConfirmed).To(BeFalse()) + }) + }) +}) diff --git a/core/services/nodes/model_load_job.go b/core/services/nodes/model_load_job.go index 14cba8a7e..843141041 100644 --- a/core/services/nodes/model_load_job.go +++ b/core/services/nodes/model_load_job.go @@ -10,6 +10,7 @@ import ( "github.com/google/uuid" "github.com/mudler/LocalAI/core/services/advisorylock" "gorm.io/gorm" + "gorm.io/gorm/clause" ) // Cold-load job states. `pending` covers node selection and replica @@ -30,17 +31,30 @@ const ( // per second regardless of how many 32 KB chunks land in it. loadJobHeartbeatInterval = stagingBroadcastInterval - // loadJobOrphanWindow is how long a job may go without a heartbeat before - // another replica may reclaim it. Generous relative to the 1s heartbeat: a - // frontend under GC pressure or a stalled DB write must not have its - // perfectly healthy multi-GB transfer stolen and restarted from zero. - loadJobOrphanWindow = 60 * time.Second + // loadJobLeaseTTL is how long a job's lease lasts after each renewal. The + // owner renews on every heartbeat, so it has about thirty chances. A frontend + // under GC pressure or a slow database write must not have a healthy + // multi-GB transfer taken from it, and a dead one must not hold a model for + // long. + loadJobLeaseTTL = 30 * time.Second - // loadJobFailureGrace is how long a failed job row is kept before deletion. - // Without it a waiter polling just after the failure finds no row, concludes - // "not loading", and starts a duplicate load of a model that just failed — - // a retry storm dressed as recovery. - loadJobFailureGrace = 15 * time.Second + // loadJobStopWindow is how long a failed job keeps its model when remote + // work may still run: the lease, a worker-side bound, and a margin. After it + // the job is released. A silent worker therefore delays a retry by a bounded + // time and never blocks the model for ever. + loadJobStopWindow = 150 * time.Second + + // loadJobLegacyStopWindow is the stop window for a worker that cannot confirm + // a stop or watch operations: the longest a load RPC may run. Such a worker + // gives no sooner bound, so the model is held no longer than the load itself + // could have run, which is what a stuck load cost before leases existed. + loadJobLegacyStopWindow = 45 * time.Minute + + // loadJobFailureReport is how long a failure is kept when the work is known + // to have ended. It exists so waiters and callers that arrive right after + // the failure read the real cause instead of starting a duplicate load of a + // model that just failed. + loadJobFailureReport = 15 * time.Second // loadJobPollInterval is how often a waiter on a non-owning replica polls // the job row. The DB is the authority: NATS staging broadcasts are @@ -66,12 +80,6 @@ func ReplicaID() string { return replicaIDValue } -// IsOrphaned reports whether the job's owner has stopped heartbeating and the -// job may be reclaimed by another replica. -func (j *ModelLoadJob) IsOrphaned(now time.Time) bool { - return now.Sub(j.LastProgress) > loadJobOrphanWindow -} - // Progress returns overall completion as a percentage, or 0 when the job has // not reported enough to compute one. func (j *ModelLoadJob) Progress() float64 { @@ -107,6 +115,181 @@ func (j *ModelLoadJob) ETA(now time.Time) (time.Duration, bool) { return time.Duration(float64(j.TotalBytes-j.BytesSent)/rate) * time.Second, true } +// LoadJobRef names one attempt to load a model. TrackingKey alone is not +// enough: a replica may delete a job and another may claim the same model, and +// a slow writer from the first attempt must not act on the second. +type LoadJobRef struct{ TrackingKey, Generation string } + +// Ref returns the attempt this row records. +func (j *ModelLoadJob) Ref() LoadJobRef { return LoadJobRef{j.TrackingKey, j.Generation} } + +// ErrStaleLoadJob means the attempt no longer owns its job: the row is gone, +// failed, or now belongs to another generation. An owner that sees it must stop. +var ErrStaleLoadJob = errors.New("stale model load job ownership") + +// ErrLoadLeaseExpired means the owner could not extend its lease before the +// lease it last held ran out. The owner stops its own work: it must not outlive +// a lease it can no longer extend. +var ErrLoadLeaseExpired = errors.New("model load job lease expired") + +// ErrLoadOperationLost means the worker no longer knows the load's operation: it +// restarted, or its watchdog ended it. The work is gone and the load cannot +// finish. +var ErrLoadOperationLost = errors.New("the worker lost the load operation") + +// loadJobResult turns the outcome of a write on one job row into an error. +// Zero rows is not a database failure: it is the proof that this attempt lost +// the job. +func loadJobResult(res *gorm.DB) error { + if res.Error != nil { + return fmt.Errorf("writing model load job: %w", res.Error) + } + if res.RowsAffected != 1 { + return ErrStaleLoadJob + } + return nil +} + +// ownedLoadJobOn is the only place a query on model_load_jobs starts. Every +// update and delete of a job goes through it, so every write carries the +// generation predicate. +func ownedLoadJobOn(db *gorm.DB, ref LoadJobRef) *gorm.DB { + return db.Model(&ModelLoadJob{}). + Where("tracking_key = ? AND generation = ?", ref.TrackingKey, ref.Generation) +} + +func (r *NodeRegistry) ownedLoadJob(ctx context.Context, ref LoadJobRef) *gorm.DB { + return ownedLoadJobOn(r.db.WithContext(ctx), ref) +} + +// now is the frontend clock. It stamps display fields only. Whether a lease +// has expired is always decided by the database clock (dbNow), so a wrong +// frontend clock cannot expire a live lease or keep a dead one. +func (r *NodeRegistry) now() time.Time { + if r.clock != nil { + return r.clock() + } + return time.Now() +} + +// dbNow is the clock every lease and deadline comparison uses. On PostgreSQL +// it is the database's own now(). SQLite has no clock of its own to speak of: it +// serves one process, so that process's clock is the database clock. +func (r *NodeRegistry) dbNow() clause.Expr { + if r.db.Dialector.Name() == "postgres" { + return gorm.Expr("now()") + } + return gorm.Expr("?", time.Now()) +} + +// dbAfter is a deadline d after dbNow. +func (r *NodeRegistry) dbAfter(d time.Duration) clause.Expr { + if r.db.Dialector.Name() == "postgres" { + return gorm.Expr("now() + make_interval(secs => ?)", d.Seconds()) + } + return gorm.Expr("?", time.Now().Add(d)) +} + +func (r *NodeRegistry) leaseExpr() clause.Expr { + ttl := loadJobLeaseTTL + if r.leaseTTL > 0 { + ttl = r.leaseTTL + } + return r.dbAfter(ttl) +} + +// activeLoadJob limits a write to an attempt that has not failed. A failed row +// only changes through its own grace-gated release. +func activeLoadJob(q *gorm.DB) *gorm.DB { + return q.Where("state <> ?", LoadJobStateFailed) +} + +// A ref with no generation names a legacy row. Nothing that runs today can own +// one, so every owner write rejects it before it reaches the database. +func (ref LoadJobRef) owned() bool { return ref.Generation != "" } + +// backfillLoadJobGenerations gives every row written before the generation +// column existed a generation of its own. The lease and stop deadline are left +// empty on purpose. A running row with no lease counts as expired, because no +// old binary renews one, and a failed row with no stop deadline gets its window +// the first time a claim or a sweep sees it. Rows an old binary writes after +// this runs follow the same path. +// +// The UUIDs are generated here, not by the database, so the migration runs the +// same on PostgreSQL and SQLite. Each write names the empty generation it +// observed, so a row another frontend already upgraded is left alone. +func backfillLoadJobGenerations(ctx context.Context, db *gorm.DB) error { + const batch = 100 + for { + var legacy []ModelLoadJob + if err := db.WithContext(ctx).Where("generation = ?", "").Limit(batch).Find(&legacy).Error; err != nil { + return err + } + if len(legacy) == 0 { + return nil + } + for _, row := range legacy { + if err := ownedLoadJobOn(db.WithContext(ctx), row.Ref()).Update("generation", uuid.NewString()).Error; err != nil { + return err + } + } + } +} + +type ( + loadOwnershipKey struct{} + loadPathKey struct{} +) + +// ErrLoadOwnershipMissing means a write on the load path carried no generation. +// The write is refused: a load that cannot say which attempt it belongs to must +// not publish anything. +var ErrLoadOwnershipMissing = errors.New("load path write without load job ownership") + +// withLoadOwnership attaches the attempt a load runs for. Registry writes made +// with the returned context are fenced on that attempt. +func withLoadOwnership(ctx context.Context, ref LoadJobRef) context.Context { + return context.WithValue(ctx, loadOwnershipKey{}, ref) +} + +// withLoadPath marks a context as running a distributed cold load. On that path +// a missing ownership value is a bug, not a caller that has nothing to fence. +func withLoadPath(ctx context.Context) context.Context { + return context.WithValue(ctx, loadPathKey{}, true) +} + +// requireLoadOwnership checks, inside the transaction that publishes a replica, +// that the attempt still owns its job. The no-op update holds the job row lock +// until the publish commits, so the owner cannot lose the job between the check +// and the write. A load-path context with no ownership fails closed. A context +// that is not on the load path (routing, health checks, tests) has nothing to +// fence and passes. +func requireLoadOwnership(ctx context.Context, tx *gorm.DB) error { + return requireLoadOwnershipFor(ctx, tx, false) +} + +// requireLoadOwnershipFor is requireLoadOwnership with a choice about a failed +// attempt. Publishing needs a live attempt. Removing the attempt's own replica +// row does not: a cancelled or failed attempt must still clean up after itself, +// and the generation still stops it from touching a successor's row. +func requireLoadOwnershipFor(ctx context.Context, tx *gorm.DB, allowFailed bool) error { + ref, ok := ctx.Value(loadOwnershipKey{}).(LoadJobRef) + if !ok { + if ctx.Value(loadPathKey{}) != nil { + return ErrLoadOwnershipMissing + } + return nil + } + if !ref.owned() { + return ErrStaleLoadJob + } + q := ownedLoadJobOn(tx, ref) + if !allowFailed { + q = activeLoadJob(q) + } + return loadJobResult(q.UpdateColumn("generation", ref.Generation)) +} + // LoadJobUpdate is a partial update to a running job. Empty node fields are // left untouched so a heartbeat does not erase the placement the runner // reported earlier. @@ -122,13 +305,30 @@ type LoadJobUpdate struct { // StartedAt anchors the rate the ETA is derived from. Set by the runner the // first time the transfer reports bytes; zero leaves the stored value alone. StartedAt time.Time + // LegacyWorker records that the worker cannot name operations. + LegacyWorker bool +} + +// claimRow is a job row together with the database's verdict on its deadlines. +type claimRow struct { + ModelLoadJob + LeaseExpired bool + StopPassed bool } // ClaimLoadJob decides, under the per-model advisory lock, whether this replica -// owns the cold load of trackingKey. It returns the live job and claimed=false -// when another replica is already loading it (or it just failed and is inside -// its grace window), or a fresh `pending` job with claimed=true when this -// replica took the work. +// owns the cold load of trackingKey. It returns claimed=true with a fresh +// `pending` job when this replica took the work, and claimed=false with the +// existing job otherwise. The rules, all judged by the database clock: +// +// 1. No row: insert. +// 2. Running, lease live: return it. The caller waits. +// 3. Running, lease expired: mark it failed with a stop window and return it. +// The caller retries once the window ends. +// 4. Failed, stop window over: delete it and insert a new generation. +// 5. Failed, window still open: return it with its cause. +// +// A row with no lease (written by an older binary) counts as expired. // // The lock is held only across these statements — no network, file, or gRPC I/O // happens inside it, which is the entire point of the job row. The primary key @@ -141,29 +341,49 @@ func (r *NodeRegistry) ClaimLoadJob(ctx context.Context, trackingKey, owner stri ) lockKey := advisorylock.KeyFromString(loadJobLockPrefix + trackingKey) err := advisorylock.WithLockCtx(ctx, r.db, lockKey, func() error { - var existing ModelLoadJob - err := r.db.WithContext(ctx).First(&existing, "tracking_key = ?", trackingKey).Error - switch { - case err == nil: - if !existing.IsOrphaned(time.Now()) { + row, found, err := r.readClaimRow(ctx, trackingKey) + if err != nil { + return err + } + if found { + existing := row.ModelLoadJob + switch { + case existing.State == LoadJobStateFailed && existing.StopDeadline == nil: + // Written by an older binary: give it the window it never got. + if err := r.setStopWindow(ctx, existing.Ref()); err != nil { + return err + } + job, claimed = r.rereadOr(ctx, trackingKey, &existing), false + return nil + case existing.State == LoadJobStateFailed && !row.StopPassed: job, claimed = &existing, false return nil + case existing.State != LoadJobStateFailed && !row.LeaseExpired: + job, claimed = &existing, false + return nil + case existing.State != LoadJobStateFailed: + // The owner stopped renewing. Fail it instead of replacing it: + // its remote work may still run, so the model stays held for + // the stop window. + err := r.expireLoadJob(ctx, existing.Ref()) + if err != nil && !errors.Is(err, ErrStaleLoadJob) { + return err + } + job, claimed = r.rereadOr(ctx, trackingKey, &existing), false + return nil } - // The owning replica died mid-load. Without this a crashed frontend - // would wedge the model permanently: every later request would find - // a job row that nobody is running and wait for a load that will - // never progress. - if err := r.db.WithContext(ctx).Delete(&ModelLoadJob{}, "tracking_key = ?", trackingKey).Error; err != nil { - return fmt.Errorf("deleting orphaned model load job: %w", err) + // Failed, and the stop window is over. The delete names the + // generation it observed, so a row another writer replaced + // meanwhile is not removed. + if err := r.ownedLoadJob(ctx, existing.Ref()).Delete(&ModelLoadJob{}).Error; err != nil { + return fmt.Errorf("deleting released model load job: %w", err) } - case errors.Is(err, gorm.ErrRecordNotFound): - default: - return fmt.Errorf("reading model load job: %w", err) } - now := time.Now() + now := r.now() fresh := &ModelLoadJob{ TrackingKey: trackingKey, + Generation: uuid.NewString(), State: LoadJobStatePending, OwnerReplica: owner, CreatedAt: now, @@ -173,6 +393,9 @@ func (r *NodeRegistry) ClaimLoadJob(ctx context.Context, trackingKey, owner stri if err := r.db.WithContext(ctx).Create(fresh).Error; err != nil { return fmt.Errorf("creating model load job: %w", err) } + if err := r.ownedLoadJob(ctx, fresh.Ref()).Update("lease_until", r.leaseExpr()).Error; err != nil { + return fmt.Errorf("leasing model load job: %w", err) + } job, claimed = fresh, true return nil }) @@ -182,6 +405,97 @@ func (r *NodeRegistry) ClaimLoadJob(ctx context.Context, trackingKey, owner stri return job, claimed, nil } +func (r *NodeRegistry) readClaimRow(ctx context.Context, trackingKey string) (claimRow, bool, error) { + var row claimRow + res := r.db.WithContext(ctx).Raw(`SELECT *, + (lease_until IS NULL OR lease_until < ?) AS lease_expired, + (stop_deadline IS NOT NULL AND stop_deadline < ?) AS stop_passed + FROM model_load_jobs WHERE tracking_key = ?`, r.dbNow(), r.dbNow(), trackingKey).Scan(&row) + if res.Error != nil { + return row, false, fmt.Errorf("reading model load job: %w", res.Error) + } + return row, res.RowsAffected > 0, nil +} + +// rereadOr returns the row as it is now, or fallback when it cannot be read. +func (r *NodeRegistry) rereadOr(ctx context.Context, trackingKey string, fallback *ModelLoadJob) *ModelLoadJob { + if job, err := r.GetLoadJob(ctx, trackingKey); err == nil && job != nil { + return job + } + return fallback +} + +// expireLoadJob fails a running job whose lease has run out, as judged by the +// database clock at the moment of the write. A renewal that landed first makes +// it a no-op and returns ErrStaleLoadJob. +func (r *NodeRegistry) expireLoadJob(ctx context.Context, ref LoadJobRef) error { + now := r.now() + return loadJobResult(activeLoadJob(r.ownedLoadJob(ctx, ref)). + Where("lease_until IS NULL OR lease_until < ?", r.dbNow()). + Updates(map[string]any{ + "state": LoadJobStateFailed, + "last_error": "the load owner stopped renewing its lease", + "op_confirmed": false, + "stop_deadline": r.dbAfter(loadJobStopWindow), + "last_progress": now, + "updated_at": now, + })) +} + +func (r *NodeRegistry) setStopWindow(ctx context.Context, ref LoadJobRef) error { + return loadJobResult(r.ownedLoadJob(ctx, ref). + Where("state = ? AND stop_deadline IS NULL", LoadJobStateFailed). + Update("stop_deadline", r.dbAfter(loadJobStopWindow))) +} + +// LoadJobSweep counts what one SweepLoadJobs pass did. +type LoadJobSweep struct { + Expired int + Released []LoadJobRef +} + +// SweepLoadJobs applies the lease rules to every job without waiting for a +// request: it fails jobs whose lease ran out and releases failed jobs whose +// stop window is over. Each write is fenced by the generation read, so a job +// replaced during the sweep is not touched. +func (r *NodeRegistry) SweepLoadJobs(ctx context.Context) (LoadJobSweep, error) { + var out LoadJobSweep + var running, failed []ModelLoadJob + if err := r.db.WithContext(ctx). + Where("state <> ? AND (lease_until IS NULL OR lease_until < ?)", LoadJobStateFailed, r.dbNow()). + Find(&running).Error; err != nil { + return out, fmt.Errorf("listing expired model load jobs: %w", err) + } + if err := r.db.WithContext(ctx). + Where("state = ? AND (stop_deadline IS NULL OR stop_deadline < ?)", LoadJobStateFailed, r.dbNow()). + Find(&failed).Error; err != nil { + return out, fmt.Errorf("listing releasable model load jobs: %w", err) + } + for _, j := range running { + switch err := r.expireLoadJob(ctx, j.Ref()); { + case err == nil: + out.Expired++ + case !errors.Is(err, ErrStaleLoadJob): + return out, err + } + } + for _, j := range failed { + if j.StopDeadline == nil { + if err := r.setStopWindow(ctx, j.Ref()); err != nil && !errors.Is(err, ErrStaleLoadJob) { + return out, err + } + continue + } + switch err := r.DeleteFailedLoadJob(ctx, j.Ref()); { + case err == nil: + out.Released = append(out.Released, j.Ref()) + case !errors.Is(err, ErrStaleLoadJob): + return out, err + } + } + return out, nil +} + // GetLoadJob returns the active job for trackingKey, or (nil, nil) when none is // active. Callers on a non-owning replica poll this; it is the authority for // both readiness and failure. @@ -206,12 +520,17 @@ func (r *NodeRegistry) ListActiveLoadJobs(ctx context.Context) ([]ModelLoadJob, return jobs, nil } -// UpdateLoadJob applies a phase transition or heartbeat. LastProgress is always -// touched: it is the liveness signal the orphan check reads, and it must tick -// even during phases that move no bytes at all. -func (r *NodeRegistry) UpdateLoadJob(ctx context.Context, trackingKey string, u LoadJobUpdate) error { - now := time.Now() +// UpdateLoadJob applies a phase transition or heartbeat and renews the lease +// from the database clock. It must tick even during phases that move no bytes +// at all: a checkpoint load moves none for many minutes. It returns ErrStaleLoadJob when +// the attempt no longer owns the job, which is the owner's signal to stop. +func (r *NodeRegistry) UpdateLoadJob(ctx context.Context, ref LoadJobRef, u LoadJobUpdate) error { + if !ref.owned() { + return ErrStaleLoadJob + } + now := r.now() fields := map[string]any{ + "lease_until": r.leaseExpr(), "last_progress": now, "updated_at": now, "bytes_sent": u.BytesSent, @@ -232,40 +551,256 @@ func (r *NodeRegistry) UpdateLoadJob(ctx context.Context, trackingKey string, u if !u.StartedAt.IsZero() { fields["started_at"] = u.StartedAt } - res := r.db.WithContext(ctx).Model(&ModelLoadJob{}). - Where("tracking_key = ?", trackingKey).Updates(fields) - if res.Error != nil { - return fmt.Errorf("updating model load job: %w", res.Error) + if u.LegacyWorker { + fields["legacy_worker"] = true } - return nil + return loadJobResult(activeLoadJob(r.ownedLoadJob(ctx, ref)).Updates(fields)) } -// FailLoadJob records the real failure on the job row so every waiter — local -// or on another replica — reports the same cause instead of an anonymous -// timeout. The row is deleted after loadJobFailureGrace by the runner. -func (r *NodeRegistry) FailLoadJob(ctx context.Context, trackingKey, msg string) error { - now := time.Now() - res := r.db.WithContext(ctx).Model(&ModelLoadJob{}). - Where("tracking_key = ?", trackingKey). - Updates(map[string]any{ - "state": LoadJobStateFailed, - "last_error": msg, - "last_progress": now, - "updated_at": now, - }) - if res.Error != nil { - return fmt.Errorf("failing model load job: %w", res.Error) +// FailLoadJob records the real failure on the job row so every waiter, local +// or on another replica, reports the same cause instead of an anonymous +// timeout. workMayRun says whether remote work may outlive the failure (a +// deadline, a cancel, a lost lease). If it may, the model stays held for the +// stop window. If the work is known to have ended, the row is kept only for the +// short report window. The next claim after that deadline replaces the row. +func (r *NodeRegistry) FailLoadJob(ctx context.Context, ref LoadJobRef, msg string, workMayRun bool) error { + if !ref.owned() { + return ErrStaleLoadJob } - return nil + hold := loadJobFailureReport + if workMayRun { + hold = loadJobStopWindow + } + now := r.now() + return loadJobResult(activeLoadJob(r.ownedLoadJob(ctx, ref)).Updates(map[string]any{ + "state": LoadJobStateFailed, + "last_error": msg, + "op_confirmed": !workMayRun, + "stop_deadline": r.dbAfter(hold), + "last_progress": now, + "updated_at": now, + })) } -// DeleteLoadJob removes a terminal job row. Success deletes immediately (the -// NodeModel row is the record of a loaded model); failures delete after their -// grace window. -func (r *NodeRegistry) DeleteLoadJob(ctx context.Context, trackingKey string) error { +// DeleteLoadJob removes the job of an attempt that succeeded. The NodeModel row +// is the record of a loaded model. A failed job is not removed here: it leaves +// through DeleteFailedLoadJob or the next claim, so a late success from a +// stale owner cannot erase the failure its waiters need to read. +func (r *NodeRegistry) DeleteLoadJob(ctx context.Context, ref LoadJobRef) error { + if !ref.owned() { + return ErrStaleLoadJob + } + return loadJobResult(activeLoadJob(r.ownedLoadJob(ctx, ref)).Delete(&ModelLoadJob{})) +} + +// DeleteFailedLoadJob releases a failed job once its stop deadline has passed +// on the database clock. The next claim does the same, so this only keeps the +// table tidy when nobody retries. +func (r *NodeRegistry) DeleteFailedLoadJob(ctx context.Context, ref LoadJobRef) error { + if !ref.owned() { + return ErrStaleLoadJob + } + return loadJobResult(r.ownedLoadJob(ctx, ref). + Where("state = ? AND stop_deadline < ?", LoadJobStateFailed, r.dbNow()). + Delete(&ModelLoadJob{})) +} + +// ConfirmLoadOp records that the remote work of a failed attempt ended: the +// worker acknowledged a stop, or it restarted. It shortens the stop window to +// the report window, so a model whose worker answers is free again in seconds +// and not after the full stop window. It never lengthens a window and never +// touches an attempt that has not failed. A replaced attempt returns +// ErrStaleLoadJob. +func (r *NodeRegistry) ConfirmLoadOp(ctx context.Context, ref LoadJobRef) error { + if !ref.owned() { + return ErrStaleLoadJob + } + if err := loadJobResult(r.ownedLoadJob(ctx, ref). + Where("state = ?", LoadJobStateFailed). + Update("op_confirmed", true)); err != nil { + return err + } + // Shorten only: a window already shorter than the report window stays. + if err := r.ownedLoadJob(ctx, ref). + Where("state = ? AND stop_deadline > ?", LoadJobStateFailed, r.dbAfter(loadJobFailureReport)). + Update("stop_deadline", r.dbAfter(loadJobFailureReport)).Error; err != nil { + return err + } + return r.removeAttemptReplicas(ctx, ref) +} + +// removeAttemptReplicas removes the replica rows a confirmed-dead attempt left +// before they reached serving. The owner's own cleanup cannot be relied on: it +// may be stuck in a call, or gone. The rows carry the attempt's generation, so +// no other attempt's row is touched. +func (r *NodeRegistry) removeAttemptReplicas(ctx context.Context, ref LoadJobRef) error { + var rows []NodeModel if err := r.db.WithContext(ctx). - Delete(&ModelLoadJob{}, "tracking_key = ?", trackingKey).Error; err != nil { - return fmt.Errorf("deleting model load job: %w", err) + Where("model_name = ? AND load_generation = ? AND state IN ?", ref.TrackingKey, ref.Generation, preServingStates). + Find(&rows).Error; err != nil { + return fmt.Errorf("listing the replicas of a dead load attempt: %w", err) + } + for _, row := range rows { + if err := r.RemoveNodeModel(ctx, row.NodeID, row.ModelName, row.ReplicaIndex); err != nil { + return err + } } return nil } + +// ConfirmNodeLoadOps confirms every failed attempt that ran on nodeID. A new +// worker incarnation calls it: the backends of the previous process exited with +// their parent, so none of that node's operations still runs. +func (r *NodeRegistry) ConfirmNodeLoadOps(ctx context.Context, nodeID string) (int, error) { + var jobs []ModelLoadJob + if err := r.db.WithContext(ctx). + Where("state = ? AND op_confirmed = ? AND node_id = ?", LoadJobStateFailed, false, nodeID). + Find(&jobs).Error; err != nil { + return 0, fmt.Errorf("listing unconfirmed load operations: %w", err) + } + confirmed := 0 + for _, j := range jobs { + switch err := r.ConfirmLoadOp(ctx, j.Ref()); { + case err == nil: + confirmed++ + case !errors.Is(err, ErrStaleLoadJob): + return confirmed, err + } + } + return confirmed, nil +} + +// ListLoadJobsAwaitingStop returns failed attempts on a known node whose remote +// work is not yet confirmed ended. The reconciler retries their stop. +func (r *NodeRegistry) ListLoadJobsAwaitingStop(ctx context.Context) ([]ModelLoadJob, error) { + var jobs []ModelLoadJob + if err := r.db.WithContext(ctx). + Where("state = ? AND op_confirmed = ? AND node_id <> ?", LoadJobStateFailed, false, ""). + Find(&jobs).Error; err != nil { + return nil, fmt.Errorf("listing load jobs awaiting stop: %w", err) + } + return jobs, nil +} + +// SetLegacyStopWindow gives a failed attempt on a worker that cannot confirm a +// stop the longest window the controller knows: the load RPC deadline. Such a +// worker does not watch operations, so nothing bounds its work sooner. +func (r *NodeRegistry) SetLegacyStopWindow(ctx context.Context, ref LoadJobRef, window time.Duration) error { + if !ref.owned() { + return ErrStaleLoadJob + } + return loadJobResult(r.ownedLoadJob(ctx, ref). + Where("state = ? AND op_confirmed = ?", LoadJobStateFailed, false). + Update("stop_deadline", r.dbAfter(window))) +} + +// CancelOutcome is what CancelLoadJob did. +type CancelOutcome int + +const ( + // CancelRecorded: the job was running and is now failed and cancelled. + CancelRecorded CancelOutcome = iota + // CancelAlready: the job had already failed. The stop window is untouched. + CancelAlready + // CancelGone: no job exists for the model. + CancelGone + // CancelConflict: another generation holds the model. The returned job is it. + CancelConflict +) + +// CancelLoadJob records an administrator's cancel of one attempt. It fails the +// attempt with the stop window, so the model is held while the remote work is +// stopped, and it marks the cancel. Repeating it never extends the window. +// The returned job is the row as it is after the call (nil for CancelGone). +func (r *NodeRegistry) CancelLoadJob(ctx context.Context, ref LoadJobRef) (CancelOutcome, *ModelLoadJob, error) { + return r.cancelLoadJob(ctx, ref, "cancelled by an administrator", true) +} + +// cancelLoadJob is CancelLoadJob with the recorded cause. byAdmin says whether +// the cancel is an administrator's request. A node that is removed or drained +// cancels its loads with a neutral cause and does not claim an administrator +// asked. +func (r *NodeRegistry) cancelLoadJob(ctx context.Context, ref LoadJobRef, reason string, byAdmin bool) (CancelOutcome, *ModelLoadJob, error) { + now := r.now() + var outcome CancelOutcome + err := advisorylock.WithLockCtx(ctx, r.db, advisorylock.KeyFromString(loadJobLockPrefix+ref.TrackingKey), func() error { + job, err := r.GetLoadJob(ctx, ref.TrackingKey) + if err != nil { + return err + } + switch { + case job == nil: + outcome = CancelGone + return nil + case job.Generation != ref.Generation: + outcome = CancelConflict + return nil + case job.State == LoadJobStateFailed: + outcome = CancelAlready + return nil + } + werr := loadJobResult(activeLoadJob(r.ownedLoadJob(ctx, ref)).Updates(map[string]any{ + "state": LoadJobStateFailed, + "cancel_requested": byAdmin, + "last_error": reason, + "op_confirmed": false, + "stop_deadline": r.dbAfter(loadJobStopWindow), + "last_progress": now, + "updated_at": now, + })) + if errors.Is(werr, ErrStaleLoadJob) { + outcome = CancelConflict + return nil + } + outcome = CancelRecorded + return werr + }) + if err != nil { + return outcome, nil, err + } + if outcome == CancelGone { + return outcome, nil, nil + } + job, err := r.GetLoadJob(ctx, ref.TrackingKey) + return outcome, job, err +} + +// ListLoadJobsOnNode returns the jobs placed on nodeID. +func (r *NodeRegistry) ListLoadJobsOnNode(ctx context.Context, nodeID string) ([]ModelLoadJob, error) { + var jobs []ModelLoadJob + if err := r.db.WithContext(ctx).Where("node_id = ?", nodeID).Find(&jobs).Error; err != nil { + return nil, fmt.Errorf("listing load jobs on node: %w", err) + } + return jobs, nil +} + +// ObserveWorkerIncarnation notes the incarnation a worker reported. When it +// differs from the one stored, the worker restarted: every operation of the +// previous process ended, so the failed attempts that ran there are confirmed. +// Most calls change nothing and cost no query, because the last value seen per +// node is cached. +func (r *NodeRegistry) ObserveWorkerIncarnation(ctx context.Context, nodeID, incarnation string) error { + if incarnation == "" { + return nil + } + if seen, ok := r.incarnations.Load(nodeID); ok && seen == incarnation { + return nil + } + var node BackendNode + if err := r.db.WithContext(ctx).Select("id", "worker_incarnation").First(&node, "id = ?", nodeID).Error; err != nil { + return nil // an unknown node is the heartbeat's business + } + if node.WorkerIncarnation != incarnation { + if err := r.db.WithContext(ctx).Model(&BackendNode{}).Where("id = ?", nodeID). + Update("worker_incarnation", incarnation).Error; err != nil { + return fmt.Errorf("recording worker incarnation: %w", err) + } + if node.WorkerIncarnation != "" { + if _, err := r.ConfirmNodeLoadOps(ctx, nodeID); err != nil { + return err + } + } + } + r.incarnations.Store(nodeID, incarnation) + return nil +} diff --git a/core/services/nodes/model_load_job_test.go b/core/services/nodes/model_load_job_test.go index 6e4d6f66a..e90e00b3b 100644 --- a/core/services/nodes/model_load_job_test.go +++ b/core/services/nodes/model_load_job_test.go @@ -74,7 +74,7 @@ var _ = Describe("ModelLoadJob", func() { // The whole point of the split: a claim is a decision that takes // milliseconds, so a concurrent request never waits behind the // minutes-long load the owner is running. - _, claimed, err := registry.ClaimLoadJob(ctx, "slow-model", "replica-a") + slow, claimed, err := registry.ClaimLoadJob(ctx, "slow-model", "replica-a") Expect(err).ToNot(HaveOccurred()) Expect(claimed).To(BeTrue()) @@ -85,7 +85,7 @@ var _ = Describe("ModelLoadJob", func() { close(running) // Stand in for a multi-GB staging run owned by replica-a. time.Sleep(2 * time.Second) - Expect(registry.DeleteLoadJob(ctx, "slow-model")).To(Succeed()) + Expect(registry.DeleteLoadJob(ctx, slow.Ref())).To(Succeed()) close(done) }() <-running @@ -99,17 +99,19 @@ var _ = Describe("ModelLoadJob", func() { <-done }) - It("reclaims a job whose owner stopped heartbeating", func() { + It("reclaims a job whose owner stopped renewing its lease", func() { _, claimed, err := registry.ClaimLoadJob(ctx, "orphan", "dead-replica") Expect(err).ToNot(HaveOccurred()) Expect(claimed).To(BeTrue()) - // Backdate the heartbeat past the orphan window, as a replica killed + // Expire the lease on the database clock, as a replica killed // mid-load would leave it. - stale := time.Now().Add(-2 * loadJobOrphanWindow) - Expect(db.Model(&ModelLoadJob{}).Where("tracking_key = ?", "orphan"). - Update("last_progress", stale).Error).ToNot(HaveOccurred()) + Expect(db.Exec("UPDATE model_load_jobs SET lease_until = now() - interval '1 second' WHERE tracking_key = 'orphan'").Error).To(Succeed()) + _, claimed, err = registry.ClaimLoadJob(ctx, "orphan", "live-replica") + Expect(err).ToNot(HaveOccurred()) + Expect(claimed).To(BeFalse(), "the dead owner's remote work may still run, so the stop window applies") + Expect(db.Exec("UPDATE model_load_jobs SET stop_deadline = now() - interval '1 second' WHERE tracking_key = 'orphan'").Error).To(Succeed()) job, claimed, err := registry.ClaimLoadJob(ctx, "orphan", "live-replica") Expect(err).ToNot(HaveOccurred()) Expect(claimed).To(BeTrue(), "a dead replica must not wedge the model permanently") @@ -134,11 +136,11 @@ var _ = Describe("ModelLoadJob", func() { }) It("records progress and clears the row on completion", func() { - _, _, err := registry.ClaimLoadJob(ctx, "m1", "replica-a") + m1Job, _, err := registry.ClaimLoadJob(ctx, "m1", "replica-a") Expect(err).ToNot(HaveOccurred()) started := time.Now() - Expect(registry.UpdateLoadJob(ctx, "m1", LoadJobUpdate{ + Expect(registry.UpdateLoadJob(ctx, m1Job.Ref(), LoadJobUpdate{ State: LoadJobStateStaging, NodeID: "node-1", NodeName: "nvidia-thor", ReplicaIndex: 2, BytesSent: 500, TotalBytes: 1000, FileIndex: 1, TotalFiles: 1, StartedAt: started, @@ -151,16 +153,16 @@ var _ = Describe("ModelLoadJob", func() { Expect(job.ReplicaIndex).To(Equal(2)) Expect(job.Progress()).To(BeNumerically("~", 50, 0.01)) - Expect(registry.DeleteLoadJob(ctx, "m1")).To(Succeed()) + Expect(registry.DeleteLoadJob(ctx, m1Job.Ref())).To(Succeed()) job, err = registry.GetLoadJob(ctx, "m1") Expect(err).ToNot(HaveOccurred()) Expect(job).To(BeNil()) }) It("keeps the placement across a byte-less heartbeat", func() { - _, _, err := registry.ClaimLoadJob(ctx, "m2", "replica-a") + m2Job, _, err := registry.ClaimLoadJob(ctx, "m2", "replica-a") Expect(err).ToNot(HaveOccurred()) - Expect(registry.UpdateLoadJob(ctx, "m2", LoadJobUpdate{ + Expect(registry.UpdateLoadJob(ctx, m2Job.Ref(), LoadJobUpdate{ State: LoadJobStateStaging, NodeID: "node-1", NodeName: "nvidia-thor", })).To(Succeed()) @@ -170,7 +172,7 @@ var _ = Describe("ModelLoadJob", func() { // A checkpoint load moves no bytes for minutes; the heartbeat must // still tick, and must not erase where the model is loading. time.Sleep(10 * time.Millisecond) - Expect(registry.UpdateLoadJob(ctx, "m2", LoadJobUpdate{State: LoadJobStateLoading})).To(Succeed()) + Expect(registry.UpdateLoadJob(ctx, m2Job.Ref(), LoadJobUpdate{State: LoadJobStateLoading})).To(Succeed()) after, err := registry.GetLoadJob(ctx, "m2") Expect(err).ToNot(HaveOccurred()) @@ -180,9 +182,9 @@ var _ = Describe("ModelLoadJob", func() { }) It("records the failure cause for waiters to read", func() { - _, _, err := registry.ClaimLoadJob(ctx, "m3", "replica-a") + m3Job, _, err := registry.ClaimLoadJob(ctx, "m3", "replica-a") Expect(err).ToNot(HaveOccurred()) - Expect(registry.FailLoadJob(ctx, "m3", "no available nodes")).To(Succeed()) + Expect(registry.FailLoadJob(ctx, m3Job.Ref(), "no available nodes", false)).To(Succeed()) job, err := registry.GetLoadJob(ctx, "m3") Expect(err).ToNot(HaveOccurred()) @@ -191,9 +193,9 @@ var _ = Describe("ModelLoadJob", func() { }) It("hands a request arriving inside the failure grace the real error", func() { - _, _, err := registry.ClaimLoadJob(ctx, "m4", "replica-a") + m4Job, _, err := registry.ClaimLoadJob(ctx, "m4", "replica-a") Expect(err).ToNot(HaveOccurred()) - Expect(registry.FailLoadJob(ctx, "m4", "worker out of VRAM")).To(Succeed()) + Expect(registry.FailLoadJob(ctx, m4Job.Ref(), "worker out of VRAM", false)).To(Succeed()) job, claimed, err := registry.ClaimLoadJob(ctx, "m4", "replica-b") Expect(err).ToNot(HaveOccurred()) diff --git a/core/services/nodes/model_loading_error.go b/core/services/nodes/model_loading_error.go index a7c8d7fb4..5c4d43ba9 100644 --- a/core/services/nodes/model_loading_error.go +++ b/core/services/nodes/model_loading_error.go @@ -2,6 +2,7 @@ package nodes import ( "fmt" + "math" "time" "github.com/mudler/LocalAI/core/schema" @@ -25,7 +26,16 @@ type ModelLoadingError struct { RetryAfter time.Duration } +// RetryLater marks the error as an answer to retry, not a fault. The model +// loader logs it quietly. +func (e *ModelLoadingError) RetryLater() bool { return true } + func (e *ModelLoadingError) Error() string { + if e.Status.State == LoadJobStateFailed { + // A model held by a failed attempt: report the real cause. The caller + // may retry once the hold ends. + return fmt.Sprintf("loading model %s: %s", e.Status.Model, e.Status.LastError) + } msg := fmt.Sprintf("model %s is %s", e.Status.Model, e.Status.State) if e.Status.Node != "" { msg += " on node " + e.Status.Node @@ -55,9 +65,35 @@ func LoadingStatus(job *ModelLoadJob) schema.ModelLoadingStatus { if eta, ok := job.ETA(time.Now()); ok { status.ETASeconds = int(eta.Seconds()) } + now := time.Now() + status.JobID = job.Generation + status.CancelRequested = job.CancelRequested + status.LastError = job.LastError + if job.LeaseUntil != nil && job.State != LoadJobStateFailed { + secs := int(job.LeaseUntil.Sub(now).Seconds()) + status.LeaseExpiresIn = &secs + } + if job.State == LoadJobStateFailed { + status.Stopping = !job.OpConfirmed + if job.StopDeadline != nil { + deadline := *job.StopDeadline + status.StopDeadline = &deadline + status.RetryAfter = max(int(math.Ceil(deadline.Sub(now).Seconds())), 0) + } + } return status } +// NewLoadHeldError is the answer for a request that finds its model held by a +// failed attempt: 503 with the real cause and a Retry-After that says when the +// hold ends. A hold always ends, so the answer is "try again", never "broken". +func NewLoadHeldError(job *ModelLoadJob) *ModelLoadingError { + status := LoadingStatus(job) + retryAfter := time.Duration(status.RetryAfter) * time.Second + retryAfter = min(max(retryAfter, time.Second), retryAfterCeiling) + return &ModelLoadingError{Status: status, RetryAfter: retryAfter} +} + // newModelLoadingError builds the 503 answer for a caller whose wait budget // expired. Retry-After is the ETA when the job has one, clamped so it stays a // useful poll interval, and the caller's own budget otherwise. diff --git a/core/services/nodes/reconciler.go b/core/services/nodes/reconciler.go index e1d7a09d3..214245351 100644 --- a/core/services/nodes/reconciler.go +++ b/core/services/nodes/reconciler.go @@ -1131,7 +1131,7 @@ func (rc *ReplicaReconciler) scaleDownIdle(ctx context.Context, cfg ModelSchedul continue } // Unload from worker - if err := rc.unloader.UnloadModelOnNode(nm.NodeID, nm.ModelName); err != nil { + if err := rc.unloader.UnloadReplica(nm.NodeID, nm); err != nil { xlog.Warn("Reconciler: unload failed (model already removed from registry)", "error", err) } xlog.Info("Reconciler: scaled down idle replica", "model", cfg.Target(), "node", nm.NodeID, "replica", nm.ReplicaIndex) diff --git a/core/services/nodes/reconciler_abandoned_load.go b/core/services/nodes/reconciler_abandoned_load.go index f8d5dc702..39a953fc3 100644 --- a/core/services/nodes/reconciler_abandoned_load.go +++ b/core/services/nodes/reconciler_abandoned_load.go @@ -26,43 +26,62 @@ const ( // answering nothing. var preServingStates = []string{"loading", "staging"} -// reclaimAbandonedLoads removes replica rows whose load will never finish. +// reclaimAbandonedLoads applies the load job lease rules, then removes replica +// rows whose load will never finish. +// +// It runs without any request. First it fails running jobs whose lease ran out +// and releases failed jobs whose stop window is over, so a crashed owner frees +// its model on its own. Then it removes replica rows stuck before serving. // // The other reconciler passes and the router's eviction query all filter // state = "loaded", and the per-model probe skips rows without an address, so -// nothing reclaimed a row that never got that far. On a node with one replica -// slot per model, a single interrupted transfer made the model unschedulable -// there until an operator intervened: scheduling saw no free slot, and eviction -// found nothing it was allowed to evict. +// nothing else reclaims a row that never got that far. On a node with one +// replica slot per model, a single interrupted transfer made the model +// unschedulable there until an operator intervened. // -// A row is only reclaimed when something proves the load is not progressing: -// either a load job that has failed or stopped heartbeating, or, for a row with -// no job at all, a node that is no longer healthy. -// -// The no-job case has to be conservative. Only the request path creates load -// jobs; the reconciler's own scale-up loads a replica without one. Treating a -// missing job as proof of abandonment would let this sweeper delete a healthy -// reconciler-driven transfer the moment it ran past the grace period, which for -// a multi-gigabyte checkpoint is every time. A healthy node with no job is -// therefore left alone; when the node is gone, nothing can be progressing and -// the row is safe to reclaim. +// A replica row names the load attempt that created it. It is removed once that +// attempt has no job: the job was released, or another attempt replaced it. +// While the job exists, even a failed one, the row keeps its slot, because the +// work behind it may still run. Rows written without a generation (an older +// binary, or a reconciler-driven load that predates job claims) are the +// uncertain case: they are removed when their model's job is released, or when +// their node is gone. A healthy node with no job is left alone, because nothing +// proves the load stopped. func (rc *ReplicaReconciler) reclaimAbandonedLoads(ctx context.Context) { if rc.db == nil { return } + // Failed attempts whose remote work is not confirmed ended keep their stop + // retried until the worker answers or the deadline releases them. + rc.retryLoadStops(ctx) + + sweep, err := rc.registry.SweepLoadJobs(ctx) + if err != nil { + xlog.Warn("Reconciler: failed to sweep load job leases, leaving replica slots held", "error", err) + return + } + if sweep.Expired > 0 || len(sweep.Released) > 0 { + xlog.Warn("Reconciler: applied load job lease rules", "expired", sweep.Expired, "released", len(sweep.Released)) + } + released := make(map[string]bool, len(sweep.Released)) + for _, ref := range sweep.Released { + released[ref.TrackingKey] = true + } + cutoff := time.Now().Add(-abandonedLoadGrace) var stuck []NodeModel + // The age grace only covers rows with no generation. A row that names its + // attempt needs no grace: its job exists from before the row does. if err := rc.db.WithContext(ctx). - Where("state IN ? AND updated_at < ?", preServingStates, cutoff). + Where("state IN ? AND (load_generation <> '' OR updated_at < ?)", preServingStates, cutoff). Find(&stuck).Error; err != nil { xlog.Warn("Reconciler: failed to list replicas stuck before serving", "error", err) return } - now := time.Now() for _, row := range stuck { - if !rc.loadAbandoned(ctx, row, now) { + if !rc.loadAbandoned(ctx, row, released[row.ModelName]) { continue } if err := rc.registry.RemoveNodeModel(ctx, row.NodeID, row.ModelName, row.ReplicaIndex); err != nil { @@ -80,23 +99,27 @@ func (rc *ReplicaReconciler) reclaimAbandonedLoads(ctx context.Context) { // // Every uncertain case answers false. Leaving a slot held for another pass // costs one scheduling opportunity; reclaiming a row out from under a live -// transfer restarts a multi-gigabyte load and, on a single-slot node, makes the -// model unschedulable there for as long as the retry loop runs. -func (rc *ReplicaReconciler) loadAbandoned(ctx context.Context, row NodeModel, now time.Time) bool { +// transfer restarts a multi-gigabyte load. +func (rc *ReplicaReconciler) loadAbandoned(ctx context.Context, row NodeModel, jobReleased bool) bool { job, err := rc.registry.GetLoadJob(ctx, row.ModelName) switch { case errors.Is(err, gorm.ErrRecordNotFound), err == nil && job == nil: - // No job: only the request path creates them, so this may be a healthy - // reconciler-driven load. Reclaim only once its node is gone. - return !rc.nodeHealthy(ctx, row.NodeID) + if row.LoadGeneration != "" { + // The attempt that made this row has no job any more. + return true + } + // No generation and no job: this may be a healthy load an older binary + // drives. Reclaim only once the node is gone, or the job it belonged to + // was just released. + return jobReleased || !rc.nodeHealthy(ctx, row.NodeID) case err != nil: xlog.Warn("Reconciler: cannot read load job, leaving the replica slot held", "model", row.ModelName, "error", err) return false - case job.State == LoadJobStateFailed: - return true default: - return job.IsOrphaned(now) + // A job exists. The row is abandoned only if the job is another + // attempt's. + return row.LoadGeneration != "" && job.Generation != row.LoadGeneration } } diff --git a/core/services/nodes/reconciler_abandoned_load_test.go b/core/services/nodes/reconciler_abandoned_load_test.go index 3289fc31b..8b0280401 100644 --- a/core/services/nodes/reconciler_abandoned_load_test.go +++ b/core/services/nodes/reconciler_abandoned_load_test.go @@ -21,8 +21,8 @@ import ( // the next request failed with "no replica slot ... all models busy". // // Elapsed time alone cannot decide this: staging a large checkpoint legitimately -// runs for tens of minutes. The load job's LastProgress heartbeat is the -// discriminator, the same signal job takeover already trusts. +// runs for tens of minutes. The load job's lease is the +// discriminator: a row names its attempt, and the attempt is alive while its job is. var _ = Describe("ReplicaReconciler — abandoned load sweeper", func() { var ( db *gorm.DB @@ -56,16 +56,27 @@ var _ = Describe("ReplicaReconciler — abandoned load sweeper", func() { }).Error).To(Succeed()) } - seedJob := func(model, state string, sinceProgress time.Duration) { + // seedJob writes a job row directly. leaseSecs and stopSecs are offsets from + // the database clock; a nil stop deadline means the job is still running. + seedJob := func(model, state, generation string, leaseSecs int, stopSecs *int) { Expect(db.Create(&ModelLoadJob{ TrackingKey: model, + Generation: generation, State: state, OwnerReplica: "someone", - LastProgress: time.Now().Add(-sinceProgress), - CreatedAt: time.Now().Add(-sinceProgress), - UpdatedAt: time.Now().Add(-sinceProgress), + LastProgress: time.Now(), + CreatedAt: time.Now(), + UpdatedAt: time.Now(), }).Error).To(Succeed()) + Expect(db.Exec("UPDATE model_load_jobs SET lease_until = now() + make_interval(secs => ?) WHERE tracking_key = ?", leaseSecs, model).Error).To(Succeed()) + if stopSecs != nil { + Expect(db.Exec("UPDATE model_load_jobs SET stop_deadline = now() + make_interval(secs => ?) WHERE tracking_key = ?", *stopSecs, model).Error).To(Succeed()) + } } + tag := func(model, generation string) { + Expect(db.Model(&NodeModel{}).Where("model_name = ?", model).Update("load_generation", generation).Error).To(Succeed()) + } + secs := func(n int) *int { return &n } rowExists := func(model string) bool { var count int64 @@ -73,15 +84,35 @@ var _ = Describe("ReplicaReconciler — abandoned load sweeper", func() { return count > 0 } - It("reclaims a staging row whose load job has stopped heartbeating", func() { + It("reclaims a staging row once the job of its attempt has been released", func() { seedReplica("abandoned", "staging", time.Hour) - seedJob("abandoned", LoadJobStateStaging, 30*time.Minute) + tag("abandoned", "gen-1") rc.reclaimAbandonedLoads(context.Background()) Expect(rowExists("abandoned")).To(BeFalse()) }) + It("reclaims a row whose attempt was replaced by another generation", func() { + seedReplica("replaced", "staging", time.Hour) + tag("replaced", "gen-1") + seedJob("replaced", LoadJobStateStaging, "gen-2", 30, nil) + + rc.reclaimAbandonedLoads(context.Background()) + + Expect(rowExists("replaced")).To(BeFalse()) + }) + + It("holds the slot while the failed job of its attempt waits out the stop window", func() { + seedReplica("holding", "staging", time.Hour) + tag("holding", "gen-1") + seedJob("holding", LoadJobStateFailed, "gen-1", -10, secs(100)) + + rc.reclaimAbandonedLoads(context.Background()) + + Expect(rowExists("holding")).To(BeTrue(), "remote work may still run until the stop deadline") + }) + It("reclaims a jobless row once its node is gone", func() { seedReplica("orphan", "loading", time.Hour) Expect(registry.MarkUnhealthy(context.Background(), node.ID)).To(Succeed()) @@ -91,32 +122,41 @@ var _ = Describe("ReplicaReconciler — abandoned load sweeper", func() { Expect(rowExists("orphan")).To(BeFalse()) }) - // Only the request path creates load jobs. The reconciler's own scale-up - // loads a replica without one, so treating a missing job as abandonment - // deleted healthy transfers the moment they outran the grace period, which - // for a multi-gigabyte checkpoint is every time. That is what made a replica - // appear to hop between nodes instead of finishing anywhere. - It("keeps a jobless row while its node is still healthy", func() { + // A row with no generation and no job may be a healthy load an older binary + // drives. Treating a missing job as abandonment deleted healthy transfers the + // moment they outran the grace period. That is what made a replica appear to + // hop between nodes instead of finishing anywhere. + It("keeps an untagged jobless row while its node is still healthy", func() { seedReplica("scaling-up", "staging", time.Hour) rc.reclaimAbandonedLoads(context.Background()) Expect(rowExists("scaling-up")).To(BeTrue(), - "a reconciler-driven load has no job row and must not be reclaimed for it") + "nothing proves the load stopped") }) - It("keeps a long transfer whose job is still heartbeating", func() { + It("reclaims an untagged row when the job of its model is released", func() { + seedReplica("legacy", "staging", time.Hour) + seedJob("legacy", LoadJobStateFailed, "gen-1", -10, secs(-1)) + + rc.reclaimAbandonedLoads(context.Background()) + + Expect(rowExists("legacy")).To(BeFalse()) + }) + + It("keeps a long transfer whose job still holds a live lease", func() { // The row itself is old, because staging does not touch it. Only the // job proves the transfer is alive. seedReplica("big-model", "staging", time.Hour) - seedJob("big-model", LoadJobStateStaging, time.Second) + tag("big-model", "gen-1") + seedJob("big-model", LoadJobStateStaging, "gen-1", 30, nil) rc.reclaimAbandonedLoads(context.Background()) Expect(rowExists("big-model")).To(BeTrue(), "a live transfer must never be reclaimed") }) - It("leaves a freshly created row alone while its job row is still being written", func() { + It("leaves a freshly created untagged row alone while its job row is still being written", func() { seedReplica("just-started", "loading", time.Second) rc.reclaimAbandonedLoads(context.Background()) @@ -126,6 +166,7 @@ var _ = Describe("ReplicaReconciler — abandoned load sweeper", func() { It("does not touch loaded replicas, which the other sweeps own", func() { seedReplica("serving", "loaded", time.Hour) + tag("serving", "gen-1") rc.reclaimAbandonedLoads(context.Background()) @@ -134,7 +175,8 @@ var _ = Describe("ReplicaReconciler — abandoned load sweeper", func() { It("frees the slot so the model can be scheduled on that node again", func() { seedReplica("wedged", "staging", time.Hour) - seedJob("wedged", LoadJobStateFailed, time.Minute) + tag("wedged", "gen-1") + seedJob("wedged", LoadJobStateFailed, "gen-1", -60, secs(-1)) _, err := registry.NextFreeReplicaIndex(context.Background(), node.ID, "wedged", 1) Expect(err).To(MatchError(ErrNoFreeSlot), "precondition: the stuck row holds the only slot") diff --git a/core/services/nodes/registry.go b/core/services/nodes/registry.go index f6daf36cf..f69d94368 100644 --- a/core/services/nodes/registry.go +++ b/core/services/nodes/registry.go @@ -22,15 +22,18 @@ import ( // Workers are generic — they don't have a fixed backend type. // The SmartRouter dynamically installs backends via NATS backend.install events. type BackendNode struct { - ID string `gorm:"primaryKey;size:36" json:"id"` - Name string `gorm:"uniqueIndex;size:255" json:"name"` - NodeType string `gorm:"size:32;default:backend" json:"node_type"` // backend, agent - Address string `gorm:"size:255" json:"address"` // host:port for gRPC - HTTPAddress string `gorm:"size:255" json:"http_address"` // host:port for HTTP file transfer - Status string `gorm:"size:32;default:registering" json:"status"` // registering, healthy, unhealthy, draining, pending - TokenHash string `gorm:"size:64" json:"-"` // SHA-256 of registration token - TotalVRAM uint64 `gorm:"column:total_vram" json:"total_vram"` // Total GPU VRAM in bytes - AvailableVRAM uint64 `gorm:"column:available_vram" json:"available_vram"` // Available GPU VRAM in bytes + // WorkerIncarnation identifies the worker process that last reported. A new + // value is proof that every load operation of the previous process ended. + WorkerIncarnation string `gorm:"size:36" json:"worker_incarnation,omitempty"` + ID string `gorm:"primaryKey;size:36" json:"id"` + Name string `gorm:"uniqueIndex;size:255" json:"name"` + NodeType string `gorm:"size:32;default:backend" json:"node_type"` // backend, agent + Address string `gorm:"size:255" json:"address"` // host:port for gRPC + HTTPAddress string `gorm:"size:255" json:"http_address"` // host:port for HTTP file transfer + Status string `gorm:"size:32;default:registering" json:"status"` // registering, healthy, unhealthy, draining, pending + TokenHash string `gorm:"size:64" json:"-"` // SHA-256 of registration token + TotalVRAM uint64 `gorm:"column:total_vram" json:"total_vram"` // Total GPU VRAM in bytes + AvailableVRAM uint64 `gorm:"column:available_vram" json:"available_vram"` // Available GPU VRAM in bytes // ReservedVRAM is a soft, in-tick reservation deducted by the scheduler when // it picks this node to load a model. Workers reset it back to 0 on each // heartbeat (the worker is the source of truth for actual free VRAM); the @@ -88,17 +91,17 @@ type BackendNode struct { // VRAMBudgetManuallySet marks the budget as a UI-set admin override so the // worker's re-registration value does not clobber it (mirrors // MaxReplicasPerModelManuallySet). - VRAMBudgetManuallySet bool `gorm:"column:vram_budget_manually_set;default:false" json:"vram_budget_manually_set"` + VRAMBudgetManuallySet bool `gorm:"column:vram_budget_manually_set;default:false" json:"vram_budget_manually_set"` // Version is the LocalAI build version reported by the worker at // registration. Empty for workers registered before this field existed. Version string `gorm:"column:version;size:64" json:"version,omitempty"` // Commit is the git commit hash the worker binary was built from. - Commit string `gorm:"column:commit;size:64" json:"commit,omitempty"` - APIKeyID string `gorm:"size:36" json:"-"` // auto-provisioned API key ID (for cleanup) - AuthUserID string `gorm:"size:36" json:"-"` // auto-provisioned user ID (for cleanup) - LastHeartbeat time.Time `gorm:"column:last_heartbeat" json:"last_heartbeat"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` + Commit string `gorm:"column:commit;size:64" json:"commit,omitempty"` + APIKeyID string `gorm:"size:36" json:"-"` // auto-provisioned API key ID (for cleanup) + AuthUserID string `gorm:"size:36" json:"-"` // auto-provisioned user ID (for cleanup) + LastHeartbeat time.Time `gorm:"column:last_heartbeat" json:"last_heartbeat"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` } const ( @@ -160,8 +163,11 @@ type NodeModel struct { CleanupError string `gorm:"column:cleanup_error;type:text" json:"cleanup_error,omitempty"` CleanupAttempts int `gorm:"column:cleanup_attempts;default:0" json:"cleanup_attempts,omitempty"` CleanupNextRetryAt *time.Time `gorm:"column:cleanup_next_retry_at" json:"cleanup_next_retry_at,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` + // LoadGeneration names the load attempt that created this replica row. A row + // still staging or loading whose attempt no longer has a job is abandoned. + LoadGeneration string `gorm:"column:load_generation;size:36;not null;default:''" json:"-"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` } // ModelLoadInfo is per-model load metadata kept independently of NodeModel rows @@ -329,7 +335,12 @@ type PendingBackendOp struct { // record of what is loaded, and keeping finished jobs would create a second // source of truth about it. type ModelLoadJob struct { - TrackingKey string `gorm:"primaryKey;size:255" json:"tracking_key"` + TrackingKey string `gorm:"primaryKey;size:255" json:"tracking_key"` + // Generation names one attempt to load the model. It is immutable for the + // life of the row, and every write of the row is conditional on it, so an + // owner that lost the job cannot change its successor. The empty string + // marks a row an older binary wrote; see backfillLoadJobGenerations. + Generation string `gorm:"size:36;not null;default:''" json:"-"` State string `gorm:"size:16;not null;index" json:"state"` OwnerReplica string `gorm:"size:64" json:"owner_replica"` NodeID string `gorm:"size:36" json:"node_id"` @@ -353,6 +364,23 @@ type ModelLoadJob struct { // minutes, so a reaper keyed on byte movement would reclaim a healthy job // mid-load. Byte progress is measured separately, by load_deadline.go. LastProgress time.Time `gorm:"index" json:"last_progress_at"` + // LeaseUntil is the owner's lease, in database time. The owner pushes it + // forward with every heartbeat. A running job whose lease is missing or in + // the past has no live owner. + LeaseUntil *time.Time `json:"-"` + // LegacyWorker is true when the load runs on a worker that cannot name + // operations. Its stop is by exact address, and its hold is the load + // deadline. + LegacyWorker bool `gorm:"not null;default:false" json:"-"` + // OpConfirmed is true once the remote work of a failed job is known to have + // ended: the owner saw the backend answer, the worker acknowledged a stop, or + // the worker restarted. It shortens the stop window. + OpConfirmed bool `gorm:"not null;default:false" json:"-"` + // CancelRequested marks a job an administrator cancelled. + CancelRequested bool `gorm:"not null;default:false" json:"-"` + // StopDeadline is set when the job fails: the earliest moment the model may + // be loaded again. It is database time too. + StopDeadline *time.Time `json:"-"` } // Op constants mirror the operation names used by DistributedBackendManager @@ -366,6 +394,13 @@ const ( // NodeRegistry manages backend node registration and lookup in PostgreSQL. type NodeRegistry struct { db *gorm.DB + // clock stamps display fields of load jobs. Tests replace it to prove that + // the lease never depends on it. nil means time.Now. + clock func() time.Time + // leaseTTL overrides loadJobLeaseTTL when set (tests). + leaseTTL time.Duration + // incarnations caches the last worker incarnation seen per node. + incarnations sync.Map // replicaRemovedHooks are invoked after a replica row for (modelName, nodeID) // is removed. This is the single chokepoint that lets dependent state be // invalidated no matter which removal path (router eviction, reconciler @@ -497,6 +532,14 @@ func NewNodeRegistry(db *gorm.DB) (*NodeRegistry, error) { return nil, fmt.Errorf("migrating node tables: %w", err) } + // Rows written before the generation column existed get one, so they follow + // the same reclaim rules as any other job. + if err := advisorylock.WithLockCtx(context.Background(), db, advisorylock.KeySchemaMigrate, func() error { + return backfillLoadJobGenerations(context.Background(), db) + }); err != nil { + return nil, fmt.Errorf("backfilling load job generations: %w", err) + } + // Rules written before scheduling rules could be keyed by an alias have no // stored target. They are all direct rules, so their target is their own // name, and the eviction guard needs the column filled in to match them. @@ -1094,6 +1137,8 @@ type HeartbeatUpdate struct { GPUVendor string `json:"gpu_vendor,omitempty"` CPUUsagePercent *float64 `json:"cpu_usage_percent,omitempty"` CPULoad1 *float64 `json:"cpu_load_1,omitempty"` + // WorkerIncarnation is the worker process identity. See BackendNode. + WorkerIncarnation string `json:"worker_incarnation,omitempty"` } func clampCPUUsage(usage float64) float64 { @@ -1382,14 +1427,23 @@ func (r *NodeRegistry) setNodeModelRevision(ctx context.Context, nodeID, modelNa // both create and update. This prevents overwriting the primary key on // subsequent calls for the same (node, model, replica_index). return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := requireLoadOwnership(ctx, tx); err != nil { + return err + } if err := requireCurrentRevision(tx, modelName, revision); err != nil { return err } + assign := map[string]any{"address": address, "state": state, "last_used": now, "in_flight": initialInFlight, + "config_revision": revision, "effective_options_hash": effectiveOptionsHash} + // A row written by a load owner names its attempt, so the reconciler can + // tell when that attempt is gone. + if ref, owned := ctx.Value(loadOwnershipKey{}).(LoadJobRef); owned { + assign["load_generation"] = ref.Generation + } var nm NodeModel return tx.Where("node_id = ? AND model_name = ? AND replica_index = ?", nodeID, modelName, replicaIndex). Attrs(NodeModel{ID: uuid.New().String(), NodeID: nodeID, ModelName: modelName, ReplicaIndex: replicaIndex}). - Assign(map[string]any{"address": address, "state": state, "last_used": now, "in_flight": initialInFlight, - "config_revision": revision, "effective_options_hash": effectiveOptionsHash}). + Assign(assign). FirstOrCreate(&nm).Error }) } @@ -1411,6 +1465,9 @@ func (r *NodeRegistry) setNodeModelLoadInfoRevision(ctx context.Context, nodeID, return err } return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := requireLoadOwnership(ctx, tx); err != nil { + return err + } if err := requireCurrentRevision(tx, modelName, revision); err != nil { return err } @@ -1453,6 +1510,9 @@ func (r *NodeRegistry) upsertModelLoadInfoRevision(ctx context.Context, modelNam UpdatedAt: now, } return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := requireLoadOwnership(ctx, tx); err != nil { + return err + } if err := requireCurrentRevision(tx, modelName, revision); err != nil { return err } @@ -1701,8 +1761,15 @@ func (r *NodeRegistry) RemoveClaimedModelCleanup(ctx context.Context, replica No // to keep the contract explicit (probeLoadedModels and scaleDownIdle iterate // per-row and must not orphan healthy siblings). func (r *NodeRegistry) RemoveNodeModel(ctx context.Context, nodeID, modelName string, replicaIndex int) error { - if err := r.db.WithContext(ctx).Where("node_id = ? AND model_name = ? AND replica_index = ?", nodeID, modelName, replicaIndex). - Delete(&NodeModel{}).Error; err != nil { + // A load owner removes its own replica row inside the fence, so a stale + // owner cannot delete the row of the attempt that replaced it. + if err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := requireLoadOwnershipFor(ctx, tx, true); err != nil { + return err + } + return tx.Where("node_id = ? AND model_name = ? AND replica_index = ?", nodeID, modelName, replicaIndex). + Delete(&NodeModel{}).Error + }); err != nil { return err } r.fireReplicaRemoved(modelName, nodeID, replicaIndex) diff --git a/core/services/nodes/router.go b/core/services/nodes/router.go index e858fb7b6..2ea18bc19 100644 --- a/core/services/nodes/router.go +++ b/core/services/nodes/router.go @@ -15,6 +15,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/services/nodes/prefixcache" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/distributedhdr" grpc "github.com/mudler/LocalAI/pkg/grpc" pb "github.com/mudler/LocalAI/pkg/grpc/proto" @@ -222,7 +223,12 @@ type SmartRouter struct { // identical outcome, so they share one wait instead of queueing. See // load_job_runner.go. loadWaitersMu sync.Mutex - loadWaiters map[string]chan struct{} + loadWaiters map[string]*loadWaiter + + // leaseTTL overrides loadJobLeaseTTL for the owner's own deadline (tests). + leaseTTL time.Duration + // opRenewEvery overrides loadOpRenewEvery, in heartbeat ticks (tests). + opRenewEvery int } // probeCacheTTL is how long a successful gRPC HealthCheck on a backend is @@ -278,7 +284,7 @@ func NewSmartRouter(registry ModelRouter, opts SmartRouterOptions) *SmartRouter stagingStallWindow: opts.StagingStallWindow, modelLoadAbsoluteMax: opts.ModelLoadAbsoluteMax, modelLoadWait: opts.ModelLoadWait, - loadWaiters: map[string]chan struct{}{}, + loadWaiters: map[string]*loadWaiter{}, } } @@ -350,6 +356,12 @@ func applyNodeHardwareDefaults(opts *pb.ModelOptions, node *BackendNode, backend func (r *SmartRouter) scheduleAndLoad(ctx context.Context, backendType, trackingKey, modelName string, configRevision string, modelOpts *pb.ModelOptions, parallel bool, initialInFlight int) (*scheduleLoadResult, error) { + // With a database every cold load runs under a job row. Mark the path so a + // registry write that lost its ownership value is refused, not waved through. + if r.db != nil { + ctx = withLoadPath(ctx) + } + node, backendAddr, replicaIndex, err := r.scheduleNewModel(ctx, backendType, trackingKey, modelOpts) if err != nil { return nil, fmt.Errorf("no available nodes: %w", err) @@ -413,6 +425,7 @@ func (r *SmartRouter) scheduleAndLoad(ctx context.Context, backendType, tracking } } + reportLoadAddress(ctx, backendAddr) client := r.buildClientForAddr(node, backendAddr, parallel) // Load the model on the remote node @@ -454,7 +467,9 @@ func (r *SmartRouter) scheduleAndLoad(ctx context.Context, backendType, tracking // minutes past the client timeout, and each retry stacked another // multi-GB loader process on the worker. Reap the replica we just // abandoned before handing the failure back. - if loadAbandonedOnWorker(err) { + // A load owner stops its own work through the operation stop path, + // so it does not reap here as well. + if _, owned := ctx.Value(loadOwnershipKey{}).(LoadJobRef); !owned && loadAbandonedOnWorker(err) { r.reapAbandonedLoad(node, trackingKey, replicaIndex) } return nil, fmt.Errorf("loading model %s on node %s: %w", modelName, node.Name, err) @@ -598,13 +613,33 @@ func (r *SmartRouter) ScheduleAndLoadModel(ctx context.Context, modelName string return nil, fmt.Errorf("unmarshalling stored model options for %s: %w", modelName, err) } - // initialInFlight=0: reconciler is pre-loading, not serving a request. - // scheduleAndLoad picks both the node and the replica slot internally. - result, err := r.scheduleAndLoad(ctx, backendType, modelName, modelName, revision, &modelOpts, false, 0) + // The reconciler is a load owner like a request is. It claims the job + // before any remote work, so a request for the same model waits for this + // load instead of scheduling a second copy, and a stale owner is stopped + // by the same loop. + job, claimed, err := r.registry.ClaimLoadJob(ctx, modelName, ReplicaID()) + if err != nil { + return nil, fmt.Errorf("claiming the load of model %s: %w", modelName, err) + } + if !claimed { + return nil, fmt.Errorf("model %s is already being loaded by another owner", modelName) + } + + var node *BackendNode + err = r.runLoadOwner(ctx, job.Ref(), func(ownerCtx context.Context) error { + // initialInFlight=0: reconciler is pre-loading, not serving a request. + // scheduleAndLoad picks both the node and the replica slot internally. + result, err := r.scheduleAndLoad(ownerCtx, backendType, modelName, modelName, revision, &modelOpts, false, 0) + if err != nil { + return err + } + node = result.Node + return nil + }) if err != nil { return nil, err } - return result.Node, nil + return node, nil } // RouteResult contains the routing decision. @@ -1367,11 +1402,25 @@ func (r *SmartRouter) installBackendOnNode(ctx context.Context, node *BackendNod // the whole time; here a cancelled ctx (typically the model-load ceiling) // frees the caller promptly. The shared install keeps running in the // background and still coalesces other callers via singleflight. + // A load owner starts the backend as an operation the worker bounds. The + // operation id is the load job generation, so a stop can name exactly this + // attempt. + ref, owned := ctx.Value(loadOwnershipKey{}).(LoadJobRef) resCh := r.installFlight.DoChan(key, func() (any, error) { - reply, err := r.unloader.InstallBackend(node.ID, backendType, modelID, r.galleriesJSON, "", "", "", replicaIndex, "", nil) + var reply *workerctl.BackendInstallReply + var err error + if owned { + reply, err = r.unloader.InstallBackendOp(node.ID, backendType, modelID, r.galleriesJSON, replicaIndex, "", ref.Generation, r.loadOperationDeadline(), nil) + } else { + reply, err = r.unloader.InstallBackend(node.ID, backendType, modelID, r.galleriesJSON, "", "", "", replicaIndex, "", nil) + } if err != nil { return "", err } + if owned && reply.Success && !reply.ReportsOperations { + // A worker that predates operations: it cannot confirm a stop. + markLegacyWorker(ctx) + } if !reply.Success { return "", fmt.Errorf("worker replied with error: %s", reply.Error) } @@ -2177,7 +2226,7 @@ func (r *SmartRouter) evictLRUAndFreeNodeFrom(ctx context.Context, candidateNode // Unload outside the transaction (NATS call) if r.unloader != nil { - if uerr := r.unloader.UnloadModelOnNode(lru.NodeID, lru.ModelName); uerr != nil { + if uerr := r.unloader.UnloadReplica(lru.NodeID, lru); uerr != nil { xlog.Warn("eviction unload failed (model already removed from registry)", "error", uerr) } } @@ -2204,3 +2253,12 @@ func (r *SmartRouter) evictLRUAndFreeNodeFrom(ctx context.Context, candidateNode return nil, ErrEvictionBusy } + +// loadOperationDeadline is the longest a worker may keep one load running: the +// controller's own absolute cap. Renewals are what normally end a load sooner. +func (r *SmartRouter) loadOperationDeadline() time.Duration { + if r.modelLoadAbsoluteMax > 0 { + return r.modelLoadAbsoluteMax + } + return modelLoadAbsoluteMax +} diff --git a/core/services/nodes/router_test.go b/core/services/nodes/router_test.go index a52d36bf9..e5a972050 100644 --- a/core/services/nodes/router_test.go +++ b/core/services/nodes/router_test.go @@ -9,6 +9,7 @@ import ( "sync" "time" + "github.com/google/uuid" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -184,12 +185,18 @@ func (s *fakeLoadJobStore) ClaimLoadJob(_ context.Context, trackingKey, owner st if s.jobs == nil { s.jobs = map[string]*ModelLoadJob{} } - if existing, ok := s.jobs[trackingKey]; ok && !existing.IsOrphaned(time.Now()) { - cp := *existing - return &cp, false, nil - } now := time.Now() - job := &ModelLoadJob{TrackingKey: trackingKey, State: LoadJobStatePending, OwnerReplica: owner, CreatedAt: now, UpdatedAt: now, LastProgress: now} + if existing, ok := s.jobs[trackingKey]; ok { + past := func(t *time.Time) bool { return t != nil && now.After(*t) } + reclaim := existing.State == LoadJobStateFailed && past(existing.StopDeadline) || + existing.State != LoadJobStateFailed && past(existing.LeaseUntil) + if !reclaim { + cp := *existing + return &cp, false, nil + } + } + lease := now.Add(loadJobLeaseTTL) + job := &ModelLoadJob{TrackingKey: trackingKey, Generation: uuid.NewString(), State: LoadJobStatePending, OwnerReplica: owner, CreatedAt: now, UpdatedAt: now, LastProgress: now, LeaseUntil: &lease} s.jobs[trackingKey] = job cp := *job return &cp, true, nil @@ -206,12 +213,22 @@ func (s *fakeLoadJobStore) GetLoadJob(_ context.Context, trackingKey string) (*M return &cp, nil } -func (s *fakeLoadJobStore) UpdateLoadJob(_ context.Context, trackingKey string, u LoadJobUpdate) error { +// owned returns the row ref still names, or nil when the attempt lost it. +// Callers hold s.mu. +func (s *fakeLoadJobStore) owned(ref LoadJobRef) *ModelLoadJob { + job, ok := s.jobs[ref.TrackingKey] + if !ok || !ref.owned() || job.Generation != ref.Generation { + return nil + } + return job +} + +func (s *fakeLoadJobStore) UpdateLoadJob(_ context.Context, ref LoadJobRef, u LoadJobUpdate) error { s.mu.Lock() defer s.mu.Unlock() - job, ok := s.jobs[trackingKey] - if !ok { - return nil + job := s.owned(ref) + if job == nil || job.State == LoadJobStateFailed { + return ErrStaleLoadJob } if u.State != "" { job.State = u.State @@ -229,24 +246,60 @@ func (s *fakeLoadJobStore) UpdateLoadJob(_ context.Context, trackingKey string, job.BytesSent, job.TotalBytes = u.BytesSent, u.TotalBytes job.FileIndex, job.TotalFiles = u.FileIndex, u.TotalFiles job.LastProgress = time.Now() + lease := time.Now().Add(loadJobLeaseTTL) + job.LeaseUntil = &lease return nil } -func (s *fakeLoadJobStore) FailLoadJob(_ context.Context, trackingKey, msg string) error { +func (s *fakeLoadJobStore) FailLoadJob(_ context.Context, ref LoadJobRef, msg string, workMayRun bool) error { s.mu.Lock() defer s.mu.Unlock() - if job, ok := s.jobs[trackingKey]; ok { - job.State = LoadJobStateFailed - job.LastError = msg - job.LastProgress = time.Now() + job := s.owned(ref) + if job == nil || job.State == LoadJobStateFailed { + return ErrStaleLoadJob } + job.State = LoadJobStateFailed + job.LastError = msg + job.LastProgress = time.Now() + hold := loadJobFailureReport + if workMayRun { + hold = loadJobStopWindow + } + deadline := time.Now().Add(hold) + job.StopDeadline = &deadline return nil } -func (s *fakeLoadJobStore) DeleteLoadJob(_ context.Context, trackingKey string) error { +func (s *fakeLoadJobStore) DeleteLoadJob(_ context.Context, ref LoadJobRef) error { s.mu.Lock() defer s.mu.Unlock() - delete(s.jobs, trackingKey) + job := s.owned(ref) + if job == nil || job.State == LoadJobStateFailed { + return ErrStaleLoadJob + } + delete(s.jobs, ref.TrackingKey) + return nil +} + +func (s *fakeLoadJobStore) ConfirmLoadOp(_ context.Context, ref LoadJobRef) error { + s.mu.Lock() + defer s.mu.Unlock() + job := s.owned(ref) + if job == nil || job.State != LoadJobStateFailed { + return ErrStaleLoadJob + } + job.OpConfirmed = true + return nil +} + +func (s *fakeLoadJobStore) DeleteFailedLoadJob(_ context.Context, ref LoadJobRef) error { + s.mu.Lock() + defer s.mu.Unlock() + job := s.owned(ref) + if job == nil || job.State != LoadJobStateFailed || job.StopDeadline == nil || time.Now().Before(*job.StopDeadline) { + return ErrStaleLoadJob + } + delete(s.jobs, ref.TrackingKey) return nil } @@ -545,6 +598,34 @@ func (f *fakeUnloader) InstallBackend(nodeID, backend, modelID, _, _, _, _ strin return f.installReply, f.installErr } +// The load operation verbs of the carrier seam. The default fake is a worker +// that names operations and acknowledges every stop. +func (f *fakeUnloader) InstallBackendOp(nodeID, backend, modelID, galleries string, replica int, opID, _ string, _ time.Duration, progress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { + reply, err := f.InstallBackend(nodeID, backend, modelID, galleries, "", "", "", replica, opID, progress) + if reply != nil { + withOps := *reply + withOps.ReportsOperations = true + reply = &withOps + } + return reply, err +} + +func (f *fakeUnloader) StopLoadOperation(_ context.Context, _ string, req workerctl.ModelStopRequest) (workerctl.ModelStopReply, error) { + return workerctl.ModelStopReply{Matched: true, Terminated: true, ProcessKey: req.ProcessKey}, nil +} + +func (f *fakeUnloader) OperationControl(_ string, req workerctl.OperationRequest) (*workerctl.OperationReply, error) { + return &workerctl.OperationReply{Renewed: req.Renew, Completed: req.Complete}, nil +} + +func (f *fakeUnloader) StopModelReplica(_ context.Context, _ string, replica NodeModel, _ bool) (workerctl.ModelStopReply, error) { + return workerctl.ModelStopReply{Matched: true, Terminated: true, ProcessKey: replica.ModelName}, nil +} + +func (f *fakeUnloader) UnloadReplica(nodeID string, replica NodeModel) error { + return f.UnloadModelOnNode(nodeID, replica.ModelName) +} + func (f *fakeUnloader) UpgradeBackend(nodeID, backend, _, _, _, _ string, replica int, _ string, _ func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendUpgradeReply, error) { f.mu.Lock() f.upgradeCalls = append(f.upgradeCalls, upgradeCall{nodeID, backend, replica}) diff --git a/core/services/nodes/unload_addressed_test.go b/core/services/nodes/unload_addressed_test.go new file mode 100644 index 000000000..cec07d48e --- /dev/null +++ b/core/services/nodes/unload_addressed_test.go @@ -0,0 +1,103 @@ +// SPDX-License-Identifier: MIT +package nodes + +import ( + "context" + "runtime" + "sync" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "gorm.io/gorm" + + "github.com/mudler/LocalAI/core/services/testutil" +) + +// addrUnloader records the unloads that name a replica's address. The worker +// frees exactly that process; a request that names only the model frees nothing. +type addrUnloader struct { + *fakeUnloader + mu sync.Mutex + unloads []NodeModel +} + +func (u *addrUnloader) UnloadReplica(_ string, replica NodeModel) error { + u.mu.Lock() + defer u.mu.Unlock() + u.unloads = append(u.unloads, replica) + return nil +} + +func (u *addrUnloader) addressed() []string { + u.mu.Lock() + defer u.mu.Unlock() + var out []string + for _, r := range u.unloads { + out = append(out, r.ModelName+"@"+r.Address) + } + return out +} + +// Both callers delete the replica row before they unload. The unload has to +// carry the address they read first, or the backend keeps the model in memory +// with nothing in the registry that points at it. +var _ = Describe("Unloading a replica that was just removed from the registry", func() { + var ( + db *gorm.DB + registry *NodeRegistry + ctx context.Context + unloader *addrUnloader + ) + + BeforeEach(func() { + if runtime.GOOS == "darwin" { + Skip("testcontainers requires Docker, not available on macOS CI") + } + db = testutil.SetupTestDB() + var err error + registry, err = NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + ctx = context.Background() + unloader = &addrUnloader{fakeUnloader: &fakeUnloader{}} + }) + + seed := func(name, model, addr string, idle time.Duration) *BackendNode { + node := &BackendNode{Name: name, NodeType: NodeTypeBackend, Address: name + ":50051"} + Expect(registry.Register(ctx, node, true)).To(Succeed()) + Expect(db.Create(&NodeModel{ID: name + "-" + model, NodeID: node.ID, ModelName: model, Address: addr, + State: "loaded", LastUsed: time.Now().Add(-idle), UpdatedAt: time.Now()}).Error).To(Succeed()) + return node + } + + It("sends an addressed unload on LRU eviction", func() { + node := seed("evict-node", "evicted-model", "10.0.0.9:9001", time.Hour) + router := NewSmartRouter(registry, SmartRouterOptions{DB: db, Unloader: unloader}) + + got, err := router.evictLRUAndFreeNodeFrom(ctx, nil) + + Expect(err).ToNot(HaveOccurred()) + Expect(got.ID).To(Equal(node.ID)) + Expect(unloader.addressed()).To(ConsistOf("evicted-model@10.0.0.9:9001")) + unloader.fakeUnloader.mu.Lock() + defer unloader.fakeUnloader.mu.Unlock() + Expect(unloader.fakeUnloader.unloadCalls).To(BeEmpty(), "no name-only unload") + }) + + It("sends an addressed unload on reconciler scale-down", func() { + n1 := seed("down-1", "scaled-model", "10.0.0.1:9001", 10*time.Minute) + n2 := seed("down-2", "scaled-model", "10.0.0.2:9002", 10*time.Minute) + Expect(registry.SetModelScheduling(ctx, &ModelSchedulingConfig{ModelName: "scaled-model", MinReplicas: 1, MaxReplicas: 4})).To(Succeed()) + _ = n1 + _ = n2 + rc := NewReplicaReconciler(ReplicaReconcilerOptions{Registry: registry, Unloader: unloader, DB: db, ScaleDownDelay: time.Minute}) + + rc.reconcile(ctx) + + Expect(unloader.addressed()).To(HaveLen(1), "one of two replicas goes") + Expect(unloader.addressed()[0]).To(MatchRegexp(`^scaled-model@10\.0\.0\.[12]:900[12]$`)) + var left int64 + Expect(db.Model(&NodeModel{}).Where("model_name = ?", "scaled-model").Count(&left).Error).To(Succeed()) + Expect(left).To(Equal(int64(1))) + }) +}) diff --git a/core/services/nodes/unloader.go b/core/services/nodes/unloader.go index ba6bc08b6..6682437a7 100644 --- a/core/services/nodes/unloader.go +++ b/core/services/nodes/unloader.go @@ -40,6 +40,9 @@ type NodeCommandSender interface { // PingNode reports whether the node is still subscribed on the bus. It // returns ErrNoRoute when nothing answers for the node. PingNode(nodeID string) error + // LoadOperationControl bounds, renews and stops remote load work. See its + // contract. + LoadOperationControl } // RemoteUnloaderAdapter implements NodeCommandSender and model.RemoteModelUnloader @@ -94,6 +97,29 @@ const exactModelStopTimeout = 10 * time.Second // cleanup intentionally has no backend.stop fallback: an old worker that does // not understand this request leaves the quarantine row for a later retry. func (a *RemoteUnloaderAdapter) StopModelReplica(ctx context.Context, nodeID string, replica NodeModel, force bool) (workerctl.ModelStopReply, error) { + return a.stopModelExact(ctx, nodeID, workerctl.ModelStopRequest{ + ModelName: replica.ModelName, + ProcessKey: model.BackendProcessKey(replica.ModelName, replica.ReplicaIndex), + ExpectedAddress: replica.Address, + Force: force, + ConfigRevision: replica.ConfigRevision, + }) +} + +// StopLoadOperation stops the process of one load operation, addressed by the +// operation id (the load job generation) and the process key. It is the only +// call that kills remote load work. It names the operation, never a model: the +// worker refuses unless the operation, the process key and, when given, the +// address and process instance all match its own records, so a stop of a load +// that finished cannot become an unload of the model. +func (a *RemoteUnloaderAdapter) StopLoadOperation(ctx context.Context, nodeID string, req workerctl.ModelStopRequest) (workerctl.ModelStopReply, error) { + if req.OperationID == "" { + return workerctl.ModelStopReply{}, errors.New("a load stop needs an operation id") + } + return a.stopModelExact(ctx, nodeID, req) +} + +func (a *RemoteUnloaderAdapter) stopModelExact(ctx context.Context, nodeID string, req workerctl.ModelStopRequest) (workerctl.ModelStopReply, error) { if ctx == nil { ctx = context.Background() } @@ -106,13 +132,7 @@ func (a *RemoteUnloaderAdapter) StopModelReplica(ctx context.Context, nodeID str } done := make(chan result, 1) go func() { - reply, err := controlRequestJSON[workerctl.ModelStopRequest, workerctl.ModelStopReply](a.nats, messaging.SubjectNodeModelStop(nodeID), workerctl.ModelStopRequest{ - ModelName: replica.ModelName, - ProcessKey: model.BackendProcessKey(replica.ModelName, replica.ReplicaIndex), - ExpectedAddress: replica.Address, - Force: force, - ConfigRevision: replica.ConfigRevision, - }, exactModelStopTimeout) + reply, err := controlRequestJSON[workerctl.ModelStopRequest, workerctl.ModelStopReply](a.nats, messaging.SubjectNodeModelStop(nodeID), req, exactModelStopTimeout) done <- result{reply: reply, err: err} }() @@ -127,6 +147,29 @@ func (a *RemoteUnloaderAdapter) StopModelReplica(ctx context.Context, nodeID str } } +// OperationControl renews and completes load operations on a node. A worker +// that predates operations never answers; the caller treats that as a node it +// cannot confirm stops on. +func (a *RemoteUnloaderAdapter) OperationControl(nodeID string, req workerctl.OperationRequest) (*workerctl.OperationReply, error) { + return controlRequestJSON[workerctl.OperationRequest, workerctl.OperationReply](a.nats, messaging.SubjectNodeModelOp(nodeID), req, operationControlTimeout) +} + +const operationControlTimeout = 5 * time.Second + +// InstallBackendOp is InstallBackend for a load operation: the request carries +// the operation id and the longest the load may run, so the worker can bound it. +func (a *RemoteUnloaderAdapter) InstallBackendOp(nodeID, backendType, modelID, galleriesJSON string, replicaIndex int, opID, operationID string, deadline time.Duration, onProgress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { + return a.installBackend(nodeID, workerctl.BackendInstallRequest{ + Backend: backendType, + ModelID: modelID, + BackendGalleries: galleriesJSON, + ReplicaIndex: int32(replicaIndex), + OpID: opID, + OperationID: operationID, + DeadlineMs: deadline.Milliseconds(), + }, onProgress) +} + // UnloadRemoteModel finds the node(s) hosting the given model and tells them // to stop their backend process via NATS backend.stop event. // The worker process handles a bounded Free() followed by process termination; @@ -215,14 +258,7 @@ func (a *RemoteUnloaderAdapter) InstallBackend( opID string, onProgress func(workerctl.BackendInstallProgressEvent), ) (*workerctl.BackendInstallReply, error) { - subject := messaging.SubjectNodeBackendInstall(nodeID) - xlog.Info("Sending NATS backend.install", "nodeID", nodeID, "backend", backendType, "modelID", modelID, "replica", replicaIndex, "opID", opID) - - // Subscribe to the per-op progress subject BEFORE publishing the install - // request so we don't miss early events. - sub := a.subscribeProgress(nodeID, opID, onProgress) - - reply, err := controlRequestJSON[workerctl.BackendInstallRequest, workerctl.BackendInstallReply](a.nats, subject, workerctl.BackendInstallRequest{ + return a.installBackend(nodeID, workerctl.BackendInstallRequest{ Backend: backendType, ModelID: modelID, BackendGalleries: galleriesJSON, @@ -231,7 +267,19 @@ func (a *RemoteUnloaderAdapter) InstallBackend( Alias: alias, ReplicaIndex: int32(replicaIndex), OpID: opID, - }, a.installTimeout) + }, onProgress) +} + +func (a *RemoteUnloaderAdapter) installBackend(nodeID string, req workerctl.BackendInstallRequest, onProgress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { + backendType, modelID, replicaIndex, opID := req.Backend, req.ModelID, int(req.ReplicaIndex), req.OpID + subject := messaging.SubjectNodeBackendInstall(nodeID) + xlog.Info("Sending NATS backend.install", "nodeID", nodeID, "backend", backendType, "modelID", modelID, "replica", replicaIndex, "opID", opID) + + // Subscribe to the per-op progress subject BEFORE publishing the install + // request so we don't miss early events. + sub := a.subscribeProgress(nodeID, opID, onProgress) + + reply, err := controlRequestJSON[workerctl.BackendInstallRequest, workerctl.BackendInstallReply](a.nats, subject, req, a.installTimeout) if sub != nil { if unsubscribeErr := sub.Unsubscribe(); unsubscribeErr != nil { @@ -545,13 +593,20 @@ func (a *RemoteUnloaderAdapter) dropStoppedReplicaRows(nodeID, op, backendName s } } -// UnloadModelOnNode sends a model.unload request to a specific node. -// The worker calls gRPC Free() to release GPU memory. -func (a *RemoteUnloaderAdapter) UnloadModelOnNode(nodeID, modelName string) error { - subject := messaging.SubjectNodeModelUnload(nodeID) - xlog.Info("Sending NATS model.unload", "nodeID", nodeID, "model", modelName) - - reply, err := controlRequestJSON[workerctl.ModelUnloadRequest, workerctl.ModelUnloadReply](a.nats, subject, workerctl.ModelUnloadRequest{ModelName: modelName}, 30*time.Second) +// UnloadReplica sends model.unload for one replica, naming its address. The +// worker calls gRPC Free() on exactly that process. A replica with no address +// never reached a backend, so there is nothing to free and nothing is sent: the +// worker is never asked to pick a process on its own. +// +// Callers that remove the replica row before they unload must pass the row they +// read first. Looking the address up afterwards finds nothing. +func (a *RemoteUnloaderAdapter) UnloadReplica(nodeID string, replica NodeModel) error { + if replica.Address == "" { + return nil + } + xlog.Info("Sending NATS model.unload", "nodeID", nodeID, "model", replica.ModelName, "replica", replica.ReplicaIndex) + reply, err := controlRequestJSON[workerctl.ModelUnloadRequest, workerctl.ModelUnloadReply](a.nats, messaging.SubjectNodeModelUnload(nodeID), + workerctl.ModelUnloadRequest{ModelName: replica.ModelName, Address: replica.Address}, 30*time.Second) if err != nil { return err } @@ -561,6 +616,32 @@ func (a *RemoteUnloaderAdapter) UnloadModelOnNode(nodeID, modelName string) erro return nil } +// UnloadModelOnNode unloads every replica of the model recorded on the node, +// each by address. It reads the replicas from the registry, so it only works +// while their rows still exist. A caller that deletes the rows first uses +// UnloadReplica with the row it read. +func (a *RemoteUnloaderAdapter) UnloadModelOnNode(nodeID, modelName string) error { + lister, ok := a.registry.(interface { + GetNodeModels(ctx context.Context, nodeID string) ([]NodeModel, error) + }) + if !ok { + return nil + } + replicas, err := lister.GetNodeModels(context.Background(), nodeID) + if err != nil { + return err + } + for _, replica := range replicas { + if replica.ModelName != modelName { + continue + } + if err := a.UnloadReplica(nodeID, replica); err != nil { + return err + } + } + return nil +} + // DeleteModelFiles sends model.delete to all nodes that have the model cached. // This removes model files from worker disks. func (a *RemoteUnloaderAdapter) DeleteModelFiles(modelName string) error { diff --git a/core/services/nodes/unloader_test.go b/core/services/nodes/unloader_test.go index 7efd39292..d3d13e269 100644 --- a/core/services/nodes/unloader_test.go +++ b/core/services/nodes/unloader_test.go @@ -324,6 +324,25 @@ var _ = Describe("RemoteUnloaderAdapter", func() { }) }) + Describe("UnloadReplica", func() { + It("sends model.unload naming the replica's address, with no registry lookup", func() { + mc.requestReply, _ = json.Marshal(workerctl.ModelUnloadReply{Success: true}) + + Expect(adapter.UnloadReplica("node-1", NodeModel{ModelName: "llama", ReplicaIndex: 1, Address: "127.0.0.1:5001"})).To(Succeed()) + + Expect(mc.requestCalls).To(HaveLen(1)) + Expect(mc.requestCalls[0].Subject).To(Equal(messaging.SubjectNodeModelUnload("node-1"))) + var request workerctl.ModelUnloadRequest + Expect(json.Unmarshal(mc.requestCalls[0].Data, &request)).To(Succeed()) + Expect(request).To(Equal(workerctl.ModelUnloadRequest{ModelName: "llama", Address: "127.0.0.1:5001"})) + }) + + It("sends nothing for a replica that never reached a backend address", func() { + Expect(adapter.UnloadReplica("node-1", NodeModel{ModelName: "llama"})).To(Succeed()) + Expect(mc.requestCalls).To(BeEmpty(), "the worker is never asked to guess a process") + }) + }) + Describe("StopModelReplica", func() { It("requests an acknowledged stop for the exact process", func() { mc.requestReply, _ = json.Marshal(workerctl.ModelStopReply{Matched: true, Terminated: true, ProcessKey: "llama#2"}) diff --git a/core/services/worker/control_nats.go b/core/services/worker/control_nats.go index cf38e6a02..b9e8dc7b2 100644 --- a/core/services/worker/control_nats.go +++ b/core/services/worker/control_nats.go @@ -36,6 +36,8 @@ func (n *natsControlServer) subject(v controlVerb) (string, error) { return messaging.SubjectNodeModelUnload(n.nodeID), nil case verbModelStop: return messaging.SubjectNodeModelStop(n.nodeID), nil + case verbModelOp: + return messaging.SubjectNodeModelOp(n.nodeID), nil case verbModelDelete: return messaging.SubjectNodeModelDelete(n.nodeID), nil case verbNodeStop: diff --git a/core/services/worker/control_nats_test.go b/core/services/worker/control_nats_test.go index 681dae63a..c115b8aac 100644 --- a/core/services/worker/control_nats_test.go +++ b/core/services/worker/control_nats_test.go @@ -129,7 +129,7 @@ var _ = Describe("Worker control verbs over NATS", func() { s = newLifecycleTestSupervisor(sigCh) }) - It("subscribes exactly the ten lifecycle subjects of the node", func() { + It("subscribes exactly the eleven lifecycle subjects of the node", func() { Expect(registerLifecycleForTest(s, bus)).To(Succeed()) Expect(bus.subscribed()).To(ConsistOf( messaging.SubjectNodeBackendInstall("n1"), @@ -140,6 +140,7 @@ var _ = Describe("Worker control verbs over NATS", func() { messaging.SubjectNodeModelsRunning("n1"), messaging.SubjectNodeModelUnload("n1"), messaging.SubjectNodeModelStop("n1"), + messaging.SubjectNodeModelOp("n1"), messaging.SubjectNodeModelDelete("n1"), messaging.SubjectNodeStop("n1"), )) @@ -175,7 +176,7 @@ var _ = Describe("Worker control verbs over NATS", func() { Eventually(replies).Should(Receive(Equal(want))) }, Entry("backend.list", messaging.SubjectNodeBackendList, `{"backends":null}`), - Entry("models.running", messaging.SubjectNodeModelsRunning, `{"models":[]}`), + Entry("models.running", messaging.SubjectNodeModelsRunning, `{"models":[],"reports_operations":true}`), ) It("signals shutdown on node.stop without replying, and never blocks on a repeat", func() { diff --git a/core/services/worker/control_server.go b/core/services/worker/control_server.go index c53037f6b..f7495eb10 100644 --- a/core/services/worker/control_server.go +++ b/core/services/worker/control_server.go @@ -22,6 +22,7 @@ const ( verbModelUnload controlVerb = "model.unload" verbModelStop controlVerb = "model.stop" verbModelDelete controlVerb = "model.delete" + verbModelOp controlVerb = "model.op" verbNodeStop controlVerb = "node.stop" verbFilesEnsure controlVerb = "files.ensure" verbFilesStage controlVerb = "files.stage" diff --git a/core/services/worker/lifecycle.go b/core/services/worker/lifecycle.go index e6561b9da..9e626dc90 100644 --- a/core/services/worker/lifecycle.go +++ b/core/services/worker/lifecycle.go @@ -46,6 +46,9 @@ func (s *backendSupervisor) registerLifecycleVerbs(srv controlServer) error { func() error { return srv.handle(verbModelStop, unary(decodeJSON[workerctl.ModelStopRequest], refuseModelStop, s.stopModelExactCtx)) }, + func() error { + return srv.handle(verbModelOp, unary(decodeJSON[workerctl.OperationRequest], refuseModelOp, s.serveOperations)) + }, func() error { return srv.handle(verbModelDelete, unary(decodeJSON[workerctl.ModelDeleteRequest], refuseModelDelete, s.deleteModel)) }, @@ -98,6 +101,11 @@ func refuseModelStop(err error) workerctl.ModelStopReply { return workerctl.ModelStopReply{Error: fmt.Sprintf("invalid request: %v", err)} } +func refuseModelOp(err error) workerctl.OperationReply { + xlog.Warn("Ignoring malformed control request", "verb", verbModelOp, "error", err) + return workerctl.OperationReply{} +} + func refuseModelDelete(err error) workerctl.ModelDeleteReply { xlog.Warn("Ignoring malformed control request", "verb", verbModelDelete, "error", err) return workerctl.ModelDeleteReply{Success: false, Error: "invalid request"} @@ -130,11 +138,20 @@ func (s *backendSupervisor) serveInstall(_ context.Context, req workerctl.Backen if install == nil { install = s.installBackend } + // The load is an operation the watchdog bounds, from the moment the request + // arrives: a controller that stops renewing it cannot leave a backend + // running for ever. + op := s.beginOperation(req) addr, err := install(req, req.Force, downloadCb) if err != nil { xlog.Error("Failed to install backend via NATS", "error", err) return workerctl.BackendInstallReply{Success: false, Error: err.Error()} } + instance := s.attachOperation(op) + if s.operationExpired(op) { + // The watchdog killed the backend while the install was still running. + return workerctl.BackendInstallReply{Success: false, Error: "load operation expired during install"} + } advertiseAddr := addr advAddr := s.cfg.advertiseAddr() @@ -148,7 +165,7 @@ func (s *backendSupervisor) serveInstall(_ context.Context, req workerctl.Backen advertiseAddr = net.JoinHostPort(advertiseHost, port) } } - return workerctl.BackendInstallReply{Success: true, Address: advertiseAddr} + return workerctl.BackendInstallReply{Success: true, Address: advertiseAddr, ProcessInstance: instance, ReportsOperations: true} } // serveUpgrade answers backend.upgrade: force-reinstall a backend. It is its @@ -395,20 +412,63 @@ func (s *backendSupervisor) unloadTargets(req workerctl.ModelUnloadRequest) []st // unloadModel answers model.unload: call gRPC Free() to release GPU memory // without killing the backend process. +// +// The target is the address in the request, or else every replica of the model +// the request names (see unloadTargets). There is no fallback to "any running +// backend": a request that names nothing running frees nothing. The supervisor +// lock is held only to snapshot each target and to verify it again afterwards, +// never across the Free() call. func (s *backendSupervisor) unloadModel(ctx context.Context, req workerctl.ModelUnloadRequest) workerctl.ModelUnloadReply { xlog.Info("Received NATS model.unload event", "model", req.ModelName) - for _, targetAddr := range s.unloadTargets(req) { - // Best-effort bounded gRPC Free(). A model.unload request must not - // occupy the NATS reply handler forever when a backend is wedged. - client := grpc.NewClientWithToken(targetAddr, false, nil, false, s.cfg.RegistrationToken) - freeCtx, cancel := context.WithTimeout(ctx, workerBackendFreeTimeout) - if err := client.Free(freeCtx); err != nil { - xlog.Warn("Free() failed during model.unload", "error", err, "addr", targetAddr) - } - cancel() + targets := s.unloadTargets(req) + if len(targets) == 0 { + xlog.Warn("model.unload names no running process; freeing nothing", "model", req.ModelName) + return workerctl.ModelUnloadReply{Success: true} } + var replaced []string + for _, addr := range targets { + // Snapshot the process at this address under the lock. Nothing there, or + // a different incarnation than the caller meant: it is already gone. + s.mu.Lock() + var target *backendProcess + var key string + for k, bp := range s.processes { + if bp.addr == addr && !bp.stopping { + target, key = bp, k + break + } + } + if target == nil || (req.ProcessInstance != "" && target.instance != req.ProcessInstance) { + s.mu.Unlock() + continue + } + instance := target.instance + s.mu.Unlock() + + // Best-effort bounded gRPC Free(), outside the lock. A model.unload + // request must not occupy the reply handler forever when a backend + // is wedged. + client := grpc.NewClientWithToken(addr, false, nil, false, s.cfg.RegistrationToken) + freeCtx, cancel := context.WithTimeout(ctx, workerBackendFreeTimeout) + if err := client.Free(freeCtx); err != nil { + xlog.Warn("Free() failed during model.unload", "error", err, "addr", addr) + } + cancel() + + // The process may have been replaced while Free() ran. Say so instead of + // reporting a success for a process that is not the one that was freed. + s.mu.Lock() + current, ok := s.processes[key] + if !ok || current != target || current.instance != instance { + replaced = append(replaced, addr) + } + s.mu.Unlock() + } + if len(replaced) > 0 { + return workerctl.ModelUnloadReply{Success: false, Error: "process was replaced during unload"} + } return workerctl.ModelUnloadReply{Success: true} } diff --git a/core/services/worker/models_running.go b/core/services/worker/models_running.go index 1f4782f87..5e3d7e496 100644 --- a/core/services/worker/models_running.go +++ b/core/services/worker/models_running.go @@ -49,9 +49,11 @@ func (s *backendSupervisor) runningModels() []workerctl.RunningModelInfo { continue } running = append(running, workerctl.RunningModelInfo{ - ModelID: modelID, - ReplicaIndex: replicaIndex, - Address: bp.addr, + ModelID: modelID, + ReplicaIndex: replicaIndex, + Address: bp.addr, + OperationID: bp.operationID, + ProcessInstance: bp.instance, }) } return running @@ -62,5 +64,5 @@ func (s *backendSupervisor) runningModels() []workerctl.RunningModelInfo { func (s *backendSupervisor) modelsRunning(_ context.Context, _ workerctl.ModelsRunningRequest) workerctl.ModelsRunningReply { running := s.runningModels() xlog.Debug("Answering models.running", "nodeID", s.nodeID, "count", len(running)) - return workerctl.ModelsRunningReply{Models: running} + return workerctl.ModelsRunningReply{Models: running, ReportsOperations: true} } diff --git a/core/services/worker/operations.go b/core/services/worker/operations.go new file mode 100644 index 000000000..0b09371dc --- /dev/null +++ b/core/services/worker/operations.go @@ -0,0 +1,268 @@ +package worker + +import ( + "context" + "fmt" + "time" + + "github.com/google/uuid" + "github.com/mudler/LocalAI/core/services/workerctl" + grpc "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/xlog" +) + +const ( + // defaultOperationKillTTL is how long a load operation may go without a + // renewal before the worker kills it. It is longer than the controller's + // 30 second lease, so the controller always expires its own row first. + defaultOperationKillTTL = 90 * time.Second + + // defaultOperationDeadline bounds an operation whose request carried no + // deadline: a controller older than the field. It matches the controller's + // own cap on one cold load. + defaultOperationDeadline = 45 * time.Minute + + // operationTickInterval is how often the watchdog looks for expired + // operations. + operationTickInterval = time.Second +) + +// workerIncarnation identifies this worker process. A new value after a restart +// is proof that every operation of the previous process ended, because the +// backends exit with their parent. +var workerIncarnation = uuid.NewString() + +// loadOperation is one load the worker watches. It lives from the install +// request until the controller completes it, the watchdog kills it, or a stop +// ends it. +type loadOperation struct { + id string + processKey string + // instance is the incarnation of the backend process the operation started + // or attached to. Empty until the process exists. + instance string + // anonymous marks an operation made for a request that named none. It is + // bounded by its deadline only, never by missing renewals: the controller + // that sent it does not know how to renew. + anonymous bool + deadline time.Time + lastRenew time.Time + // expired is set by the watchdog. A caller still inside the install reads + // it to learn that its backend was killed under it. + expired bool +} + +func (s *backendSupervisor) killTTL() time.Duration { + if s.opKillTTL > 0 { + return s.opKillTTL + } + return defaultOperationKillTTL +} + +// beginOperation registers the load a backend.install request belongs to. It +// returns nil for an install with no model (an admin backend install): there is +// no load to bound. +func (s *backendSupervisor) beginOperation(req workerctl.BackendInstallRequest) *loadOperation { + if req.ModelID == "" { + return nil + } + now := time.Now() + op := &loadOperation{ + id: req.OperationID, + processKey: model.BackendProcessKey(req.ModelID, int(req.ReplicaIndex)), + lastRenew: now, + deadline: now.Add(defaultOperationDeadline), + } + if req.DeadlineMs > 0 { + op.deadline = now.Add(time.Duration(req.DeadlineMs) * time.Millisecond) + } + if op.id == "" { + op.id = "anonymous-" + uuid.NewString() + op.anonymous = true + } + + s.mu.Lock() + defer s.mu.Unlock() + if s.operations == nil { + s.operations = make(map[string]*loadOperation) + } + if existing, ok := s.operations[op.id]; ok { + // A retry of the same install: keep its record, extend its lease. + existing.lastRenew = now + return existing + } + s.operations[op.id] = op + return op +} + +// attachOperation records which process instance an operation runs, and stamps +// the operation on the process so inventories can name it. +func (s *backendSupervisor) attachOperation(op *loadOperation) (instance string) { + if op == nil { + return "" + } + s.mu.Lock() + defer s.mu.Unlock() + if bp, ok := s.processes[op.processKey]; ok { + op.instance = bp.instance + bp.operationID = op.id + return bp.instance + } + return "" +} + +// endOperation forgets an operation. It is a no-op for an unknown one. +func (s *backendSupervisor) endOperation(id string) bool { + s.mu.Lock() + defer s.mu.Unlock() + op, ok := s.operations[id] + if !ok { + return false + } + delete(s.operations, id) + if bp, ok := s.processes[op.processKey]; ok && bp.operationID == id { + bp.operationID = "" + } + return true +} + +// operationExpired reports whether the watchdog killed this operation while the +// caller was still working on it. +func (s *backendSupervisor) operationExpired(op *loadOperation) bool { + if op == nil { + return false + } + s.mu.Lock() + defer s.mu.Unlock() + return op.expired +} + +// serveOperations answers model.op: renew and complete operations. +func (s *backendSupervisor) serveOperations(_ context.Context, req workerctl.OperationRequest) workerctl.OperationReply { + var reply workerctl.OperationReply + now := time.Now() + + s.mu.Lock() + for _, id := range req.Renew { + if op, ok := s.operations[id]; ok { + op.lastRenew = now + reply.Renewed = append(reply.Renewed, id) + } else { + reply.Unknown = append(reply.Unknown, id) + } + } + s.mu.Unlock() + + for _, id := range req.Complete { + if s.endOperation(id) { + reply.Completed = append(reply.Completed, id) + } else { + // Already ended is the state the caller asked for. + reply.Completed = append(reply.Completed, id) + } + } + return reply +} + +// runOperationWatchdog kills operations the controller stopped renewing, and +// those past their deadline. It returns when ctx ends. +func (s *backendSupervisor) runOperationWatchdog(ctx context.Context) { + tick := s.opTick + if tick <= 0 { + tick = operationTickInterval + } + ticker := time.NewTicker(tick) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + s.expireOperations(time.Now()) + } + } +} + +// expireOperations is one watchdog pass. An operation expires when its deadline +// passed or, unless it is anonymous, when no renewal arrived for the kill TTL. +// +// A backend that already answers READY has finished loading, so killing it +// would destroy a model that serves. That happens when a completion message was +// lost. The operation is then dropped, not killed. A process that cannot answer +// (still loading, or hung) is killed. +func (s *backendSupervisor) expireOperations(now time.Time) { + type victim struct { + op *loadOperation + key string + bp *backendProcess + } + var victims []victim + + s.mu.Lock() + for id, op := range s.operations { + late := now.After(op.deadline) + silent := !op.anonymous && now.Sub(op.lastRenew) > s.killTTL() + if !late && !silent { + continue + } + delete(s.operations, id) + op.expired = true + v := victim{op: op, key: op.processKey} + if bp, ok := s.processes[op.processKey]; ok && (op.instance == "" || bp.instance == op.instance) { + v.bp = bp + bp.operationID = "" + } + victims = append(victims, v) + } + s.mu.Unlock() + + for _, v := range victims { + if v.bp == nil { + xlog.Warn("Load operation expired before its process existed", "operation", v.op.id, "processKey", v.key) + continue + } + if s.backendReady(v.bp) { + xlog.Info("Load operation expired but the backend already serves; leaving it running", "operation", v.op.id, "processKey", v.key) + continue + } + xlog.Warn("Killing a load operation that was not renewed or passed its deadline", + "operation", v.op.id, "processKey", v.key) + if err := s.stopBackendExactBP(v.key, v.bp, true); err != nil { + xlog.Error("Failed to kill expired load operation", "operation", v.op.id, "processKey", v.key, "error", err) + } + } +} + +// backendReady reports whether the backend answers a status call with READY. +func (s *backendSupervisor) backendReady(bp *backendProcess) bool { + if s.readyFn != nil { + return s.readyFn(bp.addr) + } + client := grpc.NewClientWithToken(bp.addr, false, nil, false, s.cfg.RegistrationToken) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + st, err := client.Status(ctx) + return err == nil && st.GetState() == pb.StatusResponse_READY +} + +// checkOperationTarget decides whether a stop that names an operation may act on +// the process under req.ProcessKey. It holds s.mu. It refuses anything that is +// not exactly this operation's process: a stop for a load must never become an +// unload of a model that finished loading, or a stop of someone else's process. +func (s *backendSupervisor) checkOperationTarget(req workerctl.ModelStopRequest, bp *backendProcess) error { + op, known := s.operations[req.OperationID] + switch { + case known && op.processKey != req.ProcessKey: + return fmt.Errorf("operation %s runs process %s, not %s", req.OperationID, op.processKey, req.ProcessKey) + case known && op.instance != "" && bp.instance != op.instance: + return fmt.Errorf("operation %s ran a different process instance", req.OperationID) + case !known && bp.operationID != req.OperationID: + // The operation ended (or never was) and the process is not its own. + return fmt.Errorf("process %s does not belong to operation %s", req.ProcessKey, req.OperationID) + case req.ProcessInstance != "" && bp.instance != req.ProcessInstance: + return fmt.Errorf("process instance mismatch for %s", req.ProcessKey) + } + return nil +} diff --git a/core/services/worker/operations_orphan_test.go b/core/services/worker/operations_orphan_test.go new file mode 100644 index 000000000..c3f8aaf80 --- /dev/null +++ b/core/services/worker/operations_orphan_test.go @@ -0,0 +1,180 @@ +package worker + +import ( + "net" + "os" + "path/filepath" + "runtime" + "strconv" + "strings" + "time" + + process "github.com/mudler/go-processmanager" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// A worker that is killed (SIGKILL, OOM) cannot stop its backends. They keep +// running in their own process groups. These specs use real process groups and +// a fresh ledger, as a restarted worker has. +var _ = Describe("Backends a crashed worker left behind", func() { + killGroupAtEnd := func(grandchild int) { + DeferCleanup(func() { + if grandAlive(grandchild) { + _ = syscallKill(grandchild) + } + }) + } + + It("are killed, group and grandchild, when the next worker starts", func() { + if runtime.GOOS != "linux" { + Skip("the orphan sweep identifies a process by the start time in /proc, which only Linux has; elsewhere it kills nothing") + } + proc, grandchild := startGroupProcess() + killGroupAtEnd(grandchild) + path := filepath.Join(GinkgoT().TempDir(), "processes.json") + pid, err := strconv.Atoi(proc.CurrentPID()) + Expect(err).ToNot(HaveOccurred()) + newProcessLedger(path).add("model#0", pid) + + // The worker dies here: its in-memory state is gone, only the file is left. + restarted := newProcessLedger(path) + Expect(restarted.sweepStale()).To(Equal(1)) + + Eventually(func() bool { return grandAlive(grandchild) }, 10*time.Second, 50*time.Millisecond).Should(BeFalse(), + "an orphaned grandchild keeps GPU memory and a port") + Eventually(func() bool { return leaderGone(pid) }, 20*time.Second, 50*time.Millisecond).Should(BeTrue()) + }) + + It("kills a group whose leader already exited but whose grandchild lives on", func() { + if runtime.GOOS != "linux" { + Skip("the orphan sweep identifies a process by the start time in /proc, which only Linux has; elsewhere it kills nothing") + } + proc, grandchild := startGroupProcess() + killGroupAtEnd(grandchild) + path := filepath.Join(GinkgoT().TempDir(), "processes.json") + pid, _ := strconv.Atoi(proc.CurrentPID()) + newProcessLedger(path).add("model#0", pid) + // Only the leader dies. + Expect(syscallKill(pid)).To(Succeed()) + // A killed leader may stay a zombie until it is reaped. Dead or zombie + // both mean it no longer runs. + Eventually(func() bool { return leaderGone(pid) }, 20*time.Second, 50*time.Millisecond).Should(BeTrue()) + Expect(grandAlive(grandchild)).To(BeTrue()) + + Expect(newProcessLedger(path).sweepStale()).To(Equal(1)) + Eventually(func() bool { return grandAlive(grandchild) }, 10*time.Second, 50*time.Millisecond).Should(BeFalse()) + }) + + It("are not confused with an unrelated process that reused the pid", func() { + other := process.New(process.WithTemporaryStateDir(), process.WithName("/bin/sleep"), process.WithArgs("300")) + Expect(other.Run()).To(Succeed()) + DeferCleanup(func() { _ = other.Stop() }) + path := filepath.Join(GinkgoT().TempDir(), "processes.json") + pid, _ := strconv.Atoi(other.CurrentPID()) + ledger := newProcessLedger(path) + ledger.add("model#0", pid) + // The recorded start time no longer matches: the pid belongs to someone else. + ledger.corruptStartTimeForTest("model#0") + + Expect(newProcessLedger(path).sweepStale()).To(BeZero()) + Expect(pidAlive(other.CurrentPID())).To(BeTrue()) + }) + + It("are not swept once the worker stopped them itself", func() { + other := process.New(process.WithTemporaryStateDir(), process.WithName("/bin/sleep"), process.WithArgs("300")) + Expect(other.Run()).To(Succeed()) + DeferCleanup(func() { _ = other.Stop() }) + path := filepath.Join(GinkgoT().TempDir(), "processes.json") + pid, _ := strconv.Atoi(other.CurrentPID()) + ledger := newProcessLedger(path) + ledger.add("model#0", pid) + ledger.remove("model#0") + + Expect(newProcessLedger(path).sweepStale()).To(BeZero()) + Expect(pidAlive(other.CurrentPID())).To(BeTrue()) + }) + + It("kills nothing where start times cannot be read, as on macOS", func() { + other := process.New(process.WithTemporaryStateDir(), process.WithName("/bin/sleep"), process.WithArgs("300")) + Expect(other.Run()).To(Succeed()) + DeferCleanup(func() { _ = other.Stop() }) + path := filepath.Join(GinkgoT().TempDir(), "processes.json") + pid, _ := strconv.Atoi(other.CurrentPID()) + ledger := newProcessLedger(path) + ledger.add("model#0", pid) // recorded with a start time, as on Linux + ledger.corruptStartTimeForTest("model#0") + + real := readStartTime + readStartTime = func(int) string { return "" } // no /proc + DeferCleanup(func() { readStartTime = real }) + + Expect(ledger.sweepStale()).To(BeZero(), "with no way to tell whose it is, nothing is killed") + Expect(pidAlive(other.CurrentPID())).To(BeTrue()) + }) + + It("never kills an entry recorded without a start time", func() { + other := process.New(process.WithTemporaryStateDir(), process.WithName("/bin/sleep"), process.WithArgs("300")) + Expect(other.Run()).To(Succeed()) + DeferCleanup(func() { _ = other.Stop() }) + path := filepath.Join(GinkgoT().TempDir(), "processes.json") + pid, _ := strconv.Atoi(other.CurrentPID()) + ledger := newProcessLedger(path) + ledger.add("model#0", pid) + ledger.forgetStartTimeForTest("model#0") // as on a system with no /proc + + Expect(newProcessLedger(path).sweepStale()).To(BeZero()) + Expect(pidAlive(other.CurrentPID())).To(BeTrue()) + }) + + It("is a no-op with no ledger file", func() { + Expect(newProcessLedger(filepath.Join(GinkgoT().TempDir(), "missing.json")).sweepStale()).To(BeZero()) + var none *processLedger + Expect(none.sweepStale()).To(BeZero()) + none.add("x", 1) // must not panic + none.remove("x") + }) +}) + +// A restarted worker can hand out a port an orphan still holds. The readiness +// poll would then connect to the orphan and report a backend that is not its own. +var _ = Describe("Port allocation", func() { + It("skips a port that something already listens on", func() { + lis, err := net.Listen("tcp", "127.0.0.1:0") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = lis.Close() }) + busy := lis.Addr().(*net.TCPAddr).Port + + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{}, nextPort: busy, minPort: busy, maxPort: busy + 20} + s.mu.Lock() + port, err := s.allocateFreePort("model#0") + s.mu.Unlock() + + Expect(err).ToNot(HaveOccurred()) + Expect(port).ToNot(Equal(busy), "an orphan holds that port") + Expect(quarantinedPortNumbers(s)).To(ContainElement(busy), "the busy port comes back later, not now") + }) + + It("refuses a start when every port in range is busy", func() { + lis, err := net.Listen("tcp", "127.0.0.1:0") + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(func() { _ = lis.Close() }) + busy := lis.Addr().(*net.TCPAddr).Port + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{}, nextPort: busy, minPort: busy, maxPort: busy} + s.mu.Lock() + _, err = s.allocateFreePort("model#0") + s.mu.Unlock() + Expect(err).To(MatchError(ErrNoFreePort)) + }) +}) + +// leaderGone reports whether pid is gone or only a zombie waiting to be reaped. +func leaderGone(pid int) bool { + data, err := os.ReadFile("/proc/" + strconv.Itoa(pid) + "/stat") + if err != nil { + return true + } + text := string(data) + i := strings.LastIndex(text, ")") + return i >= 0 && len(text) > i+2 && text[i+2] == 'Z' +} diff --git a/core/services/worker/operations_test.go b/core/services/worker/operations_test.go new file mode 100644 index 000000000..afaf1352c --- /dev/null +++ b/core/services/worker/operations_test.go @@ -0,0 +1,344 @@ +package worker + +import ( + "context" + "net" + "os" + "path/filepath" + "strconv" + "strings" + "sync/atomic" + "syscall" + "time" + + process "github.com/mudler/go-processmanager" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/workerctl" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + gogrpc "google.golang.org/grpc" +) + +// freeHookBackend runs a hook inside Free(), so a spec can act while the +// supervisor is in the middle of an unload. +type freeHookBackend struct { + pb.UnimplementedBackendServer + frees atomic.Int32 + onFree func() +} + +func (b *freeHookBackend) Free(context.Context, *pb.HealthMessage) (*pb.Result, error) { + b.frees.Add(1) + if b.onFree != nil { + b.onFree() + } + return &pb.Result{Success: true}, nil +} + +// startGroupProcess starts a real shell in its own process group. The shell +// starts a grandchild and records its pid, so a spec can tell that the whole +// group is gone and not only the leader. +func startGroupProcess() (*process.Process, int) { + pidFile := filepath.Join(GinkgoT().TempDir(), "grandchild.pid") + proc := process.New(process.WithTemporaryStateDir(), process.WithName("/bin/sh"), + process.WithArgs("-c", "sleep 300 & echo $! > "+pidFile+"; wait")) + Expect(proc.Run()).To(Succeed()) + var grandchild int + Eventually(func() int { + data, err := os.ReadFile(pidFile) + if err != nil { + return 0 + } + grandchild, _ = strconv.Atoi(strings.TrimSpace(string(data))) + return grandchild + }, 5*time.Second, 20*time.Millisecond).ShouldNot(BeZero()) + return proc, grandchild +} + +func newOperationSupervisor(proc *process.Process) (*backendSupervisor, *loadOperation) { + s := &backendSupervisor{ + cfg: &Config{}, + processes: map[string]*backendProcess{}, + opKillTTL: time.Second, + opTick: 100 * time.Millisecond, + readyFn: func(string) bool { return false }, + } + s.processes["model#0"] = &backendProcess{proc: proc, addr: "127.0.0.1:59001", port: 59001, instance: "instance-1"} + op := s.beginOperation(workerctl.BackendInstallRequest{ModelID: "model", OperationID: "gen-1", DeadlineMs: 600000}) + Expect(s.attachOperation(op)).To(Equal("instance-1")) + return s, op +} + +func runWatchdog(s *backendSupervisor) { + ctx, cancel := context.WithCancel(context.Background()) + DeferCleanup(cancel) + go s.runOperationWatchdog(ctx) +} + +var _ = Describe("Load operation watchdog", func() { + It("kills the whole process group, grandchildren included, when renewals stop", func() { + proc, grandchild := startGroupProcess() + DeferCleanup(func() { + if grandAlive(grandchild) { + _ = syscallKill(grandchild) + } + }) + s, _ := newOperationSupervisor(proc) + runWatchdog(s) + + // Renewals keep it alive well past the kill TTL. + for range 12 { + time.Sleep(200 * time.Millisecond) + reply := s.serveOperations(context.Background(), workerctl.OperationRequest{Renew: []string{"gen-1"}}) + Expect(reply.Renewed).To(ConsistOf("gen-1")) + } + Expect(grandAlive(grandchild)).To(BeTrue()) + Expect(pidAlive(proc.CurrentPID())).To(BeTrue()) + + // The controller goes silent. + Eventually(func() bool { return grandAlive(grandchild) }, 15*time.Second, 100*time.Millisecond).Should(BeFalse(), + "the grandchild is in the group and must die with it") + Eventually(proc.Done(), 15*time.Second, 100*time.Millisecond).Should(BeClosed(), "the leader is reaped once it is killed") + Eventually(func() int { s.mu.Lock(); defer s.mu.Unlock(); return len(s.processes) }, 5*time.Second).Should(BeZero()) + }) + + It("kills at the absolute deadline even when renewals keep arriving", func() { + proc, grandchild := startGroupProcess() + DeferCleanup(func() { + if grandAlive(grandchild) { + _ = syscallKill(grandchild) + } + }) + s, op := newOperationSupervisor(proc) + s.mu.Lock() + op.deadline = time.Now().Add(700 * time.Millisecond) + s.mu.Unlock() + runWatchdog(s) + stop := make(chan struct{}) + defer close(stop) + go func() { + for { + select { + case <-stop: + return + case <-time.After(100 * time.Millisecond): + s.serveOperations(context.Background(), workerctl.OperationRequest{Renew: []string{"gen-1"}}) + } + } + }() + Eventually(func() bool { return grandAlive(grandchild) }, 15*time.Second, 100*time.Millisecond).Should(BeFalse()) + }) + + It("never kills an anonymous operation for missing renewals, only at its deadline", func() { + proc, grandchild := startGroupProcess() + DeferCleanup(func() { + if grandAlive(grandchild) { + _ = syscallKill(grandchild) + } + }) + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{}, + opKillTTL: 200 * time.Millisecond, opTick: 50 * time.Millisecond, readyFn: func(string) bool { return false }} + s.processes["model#0"] = &backendProcess{proc: proc, addr: "127.0.0.1:59002", port: 59002, instance: "i"} + // A controller older than operations sends no id and no deadline. + op := s.beginOperation(workerctl.BackendInstallRequest{ModelID: "model"}) + Expect(op.anonymous).To(BeTrue()) + s.attachOperation(op) + runWatchdog(s) + + Consistently(func() bool { return grandAlive(grandchild) }, 1500*time.Millisecond, 100*time.Millisecond).Should(BeTrue(), + "silence is not a reason to kill a controller that cannot renew") + + s.mu.Lock() + op.deadline = time.Now().Add(300 * time.Millisecond) + s.mu.Unlock() + Eventually(func() bool { return grandAlive(grandchild) }, 15*time.Second, 100*time.Millisecond).Should(BeFalse()) + }) + + It("does not kill a backend that already serves, when a completion was lost", func() { + proc, grandchild := startGroupProcess() + DeferCleanup(func() { + if grandAlive(grandchild) { + _ = syscallKill(grandchild) + } + }) + s, _ := newOperationSupervisor(proc) + s.readyFn = func(string) bool { return true } + runWatchdog(s) + + Eventually(func() int { s.mu.Lock(); defer s.mu.Unlock(); return len(s.operations) }, 10*time.Second).Should(BeZero()) + Expect(grandAlive(grandchild)).To(BeTrue(), "a serving model is not a load to kill") + Expect(s.runningModels()).To(HaveLen(1)) + }) + + It("stops watching an operation once it is completed", func() { + proc, grandchild := startGroupProcess() + DeferCleanup(func() { + if grandAlive(grandchild) { + _ = syscallKill(grandchild) + } + _ = proc.Stop() + }) + s, _ := newOperationSupervisor(proc) + runWatchdog(s) + + reply := s.serveOperations(context.Background(), workerctl.OperationRequest{Complete: []string{"gen-1"}}) + Expect(reply.Completed).To(ConsistOf("gen-1")) + Consistently(func() bool { return grandAlive(grandchild) }, 2*time.Second, 100*time.Millisecond).Should(BeTrue()) + Expect(s.runningModels()[0].OperationID).To(BeEmpty()) + }) + + It("tells the controller which operations it does not know", func() { + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{}} + reply := s.serveOperations(context.Background(), workerctl.OperationRequest{Renew: []string{"nope"}}) + Expect(reply.Unknown).To(ConsistOf("nope")) + Expect(reply.Renewed).To(BeEmpty()) + }) +}) + +var _ = Describe("Stopping one operation", func() { + It("reports the operation and the process instance in the inventory", func() { + proc := startModelStopProcess() + DeferCleanup(func() { _ = proc.Stop() }) + s, _ := newOperationSupervisor(proc) + reply := s.modelsRunning(context.Background(), workerctl.ModelsRunningRequest{}) + Expect(reply.ReportsOperations).To(BeTrue()) + Expect(reply.Models).To(ConsistOf(workerctl.RunningModelInfo{ + ModelID: "model", ReplicaIndex: 0, Address: "127.0.0.1:59001", OperationID: "gen-1", ProcessInstance: "instance-1", + })) + }) + + It("stops exactly the process of the named operation and forgets the operation", func() { + proc := startModelStopProcess() + s, _ := newOperationSupervisor(proc) + + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: "127.0.0.1:59001", + OperationID: "gen-1", ProcessInstance: "instance-1", Force: true}) + + Expect(reply.Error).To(BeEmpty()) + Expect(reply.Terminated).To(BeTrue()) + Expect(proc.Done()).To(BeClosed()) + Expect(s.operations).To(BeEmpty()) + }) + + It("refuses a stop whose operation is not the one the process runs", func() { + proc := startModelStopProcess() + DeferCleanup(func() { _ = proc.Stop() }) + s, _ := newOperationSupervisor(proc) + + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: "127.0.0.1:59001", + OperationID: "other-generation", Force: true}) + + Expect(reply.Terminated).To(BeFalse()) + Expect(reply.Error).To(ContainSubstring("does not belong to operation")) + Expect(pidAlive(proc.CurrentPID())).To(BeTrue()) + Expect(s.operations).To(HaveKey("gen-1")) + }) + + It("refuses to turn a stop of a finished load into an unload of the model", func() { + proc := startModelStopProcess() + DeferCleanup(func() { _ = proc.Stop() }) + s, _ := newOperationSupervisor(proc) + s.serveOperations(context.Background(), workerctl.OperationRequest{Complete: []string{"gen-1"}}) + + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: "127.0.0.1:59001", + OperationID: "gen-1", Force: true}) + + Expect(reply.Terminated).To(BeFalse(), "a cancel of a load that finished must leave the serving model alone") + Expect(pidAlive(proc.CurrentPID())).To(BeTrue()) + }) + + It("refuses a stop that names another process instance", func() { + proc := startModelStopProcess() + DeferCleanup(func() { _ = proc.Stop() }) + s, _ := newOperationSupervisor(proc) + + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: "127.0.0.1:59001", + OperationID: "gen-1", ProcessInstance: "a-replacement", Force: true}) + + Expect(reply.Terminated).To(BeFalse()) + Expect(reply.Error).To(ContainSubstring("instance")) + Expect(pidAlive(proc.CurrentPID())).To(BeTrue()) + }) + + It("answers terminated for an operation whose process is already gone", func() { + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{}} + s.beginOperation(workerctl.BackendInstallRequest{ModelID: "model", OperationID: "gen-9"}) + + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", OperationID: "gen-9"}) + + Expect(reply.Terminated).To(BeTrue()) + Expect(s.operations).To(BeEmpty()) + }) +}) + +var _ = Describe("model.unload", func() { + It("frees nothing when the request names no running model, and never another model's process", func() { + backend := &freeHookBackend{} + addr, port, stopServer := startFreeHookBackend(backend) + defer stopServer() + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{ + "model#0": {addr: addr, port: port, instance: "i"}, + }} + + Expect(s.unloadModel(context.Background(), workerctl.ModelUnloadRequest{}).Success).To(BeTrue()) + Expect(s.unloadModel(context.Background(), workerctl.ModelUnloadRequest{ModelName: "another-model"}).Success).To(BeTrue()) + Expect(backend.frees.Load()).To(BeZero(), "the worker must not guess which running backend to free") + + // A request that names the model frees that model's own process. + Expect(s.unloadModel(context.Background(), workerctl.ModelUnloadRequest{ModelName: "model"}).Success).To(BeTrue()) + Expect(backend.frees.Load()).To(Equal(int32(1))) + }) + + It("frees only the process at the given address and instance", func() { + backend := &freeHookBackend{} + addr, port, stopServer := startFreeHookBackend(backend) + defer stopServer() + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{ + "model#0": {addr: addr, port: port, instance: "i"}, + }} + + Expect(s.unloadModel(context.Background(), workerctl.ModelUnloadRequest{Address: addr, ProcessInstance: "other"}).Success).To(BeTrue()) + Expect(backend.frees.Load()).To(BeZero()) + Expect(s.unloadModel(context.Background(), workerctl.ModelUnloadRequest{Address: addr, ProcessInstance: "i"}).Success).To(BeTrue()) + Expect(backend.frees.Load()).To(Equal(int32(1))) + }) + + It("does not hold the supervisor lock across Free, and notices a replaced process", func() { + backend := &freeHookBackend{} + addr, port, stopServer := startFreeHookBackend(backend) + defer stopServer() + s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{ + "model#0": {addr: addr, port: port, instance: "i"}, + }} + // If unloadModel held the lock during Free, taking it here would + // deadlock and the spec would time out. + backend.onFree = func() { + s.mu.Lock() + s.processes["model#0"] = &backendProcess{addr: addr, port: port, instance: "replacement"} + s.mu.Unlock() + } + + reply := s.unloadModel(context.Background(), workerctl.ModelUnloadRequest{Address: addr}) + + Expect(reply.Success).To(BeFalse()) + Expect(reply.Error).To(ContainSubstring("replaced")) + }) +}) + +func grandAlive(pid int) bool { return pidAlive(strconv.Itoa(pid)) } + +func syscallKill(pid int) error { return syscall.Kill(pid, syscall.SIGKILL) } + +func startRegisteredBackend(register func(*gogrpc.Server)) (string, int, func()) { + lis, err := net.Listen("tcp", "127.0.0.1:0") + Expect(err).NotTo(HaveOccurred()) + server := gogrpc.NewServer() + register(server) + go func() { _ = server.Serve(lis) }() + return lis.Addr().String(), lis.Addr().(*net.TCPAddr).Port, server.Stop +} + +func startFreeHookBackend(backend *freeHookBackend) (string, int, func()) { + return startRegisteredBackend(func(server *gogrpc.Server) { pb.RegisterBackendServer(server, backend) }) +} diff --git a/core/services/worker/process_ledger.go b/core/services/worker/process_ledger.go new file mode 100644 index 000000000..db8b898ea --- /dev/null +++ b/core/services/worker/process_ledger.go @@ -0,0 +1,189 @@ +package worker + +import ( + "encoding/json" + "os" + "path/filepath" + "strconv" + "strings" + "sync" + + "github.com/mudler/xlog" +) + +// processLedger remembers the backend processes this worker started, in a small +// file, so the next worker can find the ones this one left behind. +// +// Backends run in their own process groups. When the worker is killed (SIGKILL, +// the OOM killer) nothing stops them: they keep their port and their GPU +// memory, and a restarted worker has no record of them. The ledger is that +// record. A restarted worker kills every group it lists before it serves. +// +// An entry is the leader's pid and its start time. The start time is what keeps +// a recycled pid from being mistaken for a backend. A nil ledger does nothing. +type processLedger struct { + path string + + mu sync.Mutex + entries map[string]ledgerEntry +} + +type ledgerEntry struct { + Key string `json:"key"` + PID int `json:"pid"` + StartTime string `json:"start_time"` +} + +func newProcessLedger(path string) *processLedger { + return &processLedger{path: path, entries: map[string]ledgerEntry{}} +} + +// add records a started process. A failure is logged and ignored: losing the +// ledger costs the orphan sweep, never a start. +func (l *processLedger) add(key string, pid int) { + if l == nil || pid <= 0 { + return + } + l.mu.Lock() + defer l.mu.Unlock() + l.entries[key] = ledgerEntry{Key: key, PID: pid, StartTime: readStartTime(pid)} + l.flushLocked() +} + +// remove forgets a process the worker stopped itself. +func (l *processLedger) remove(key string) { + if l == nil { + return + } + l.mu.Lock() + defer l.mu.Unlock() + if _, ok := l.entries[key]; !ok { + return + } + delete(l.entries, key) + l.flushLocked() +} + +// sweepStale kills the process groups a previous worker recorded and did not +// stop, and returns how many it killed. It runs once, before this worker starts +// any backend, so every entry in the file belongs to a predecessor. +func (l *processLedger) sweepStale() int { + if l == nil { + return 0 + } + l.mu.Lock() + defer l.mu.Unlock() + + data, err := os.ReadFile(l.path) + if err != nil { + return 0 + } + // Where the start time of a process cannot be read at all, no entry can be + // told apart from an unrelated process that reused its pid. Nothing is killed. + if !startTimesReadable() { + l.entries = map[string]ledgerEntry{} + l.flushLocked() + return 0 + } + var stale []ledgerEntry + if err := json.Unmarshal(data, &stale); err != nil { + xlog.Warn("Ignoring an unreadable worker process ledger", "path", l.path, "error", err) + } + killed := 0 + for _, e := range stale { + if e.PID <= 0 { + continue + } + // Without the leader's start time at the time it was recorded, there is no + // way to tell the process apart from an unrelated one that reused the pid + // (the start time is only readable on Linux). Such an entry is never + // killed. + if e.StartTime == "" { + continue + } + // The leader still exists with another start time: its pid was reused by + // an unrelated process. A group that lost its leader cannot be reused (the + // kernel keeps the pid while the group lives), so a missing leader is + // ours. + if now := readStartTime(e.PID); now != "" && now != e.StartTime { + continue + } + if err := killProcessGroup(e.PID); err == nil { + killed++ + xlog.Warn("Killed a backend process group left behind by a previous worker", "processKey", e.Key, "pid", e.PID) + } + } + l.entries = map[string]ledgerEntry{} + l.flushLocked() + return killed +} + +// corruptStartTimeForTest makes an entry look like a recycled pid. +func (l *processLedger) corruptStartTimeForTest(key string) { + l.mu.Lock() + defer l.mu.Unlock() + e := l.entries[key] + e.StartTime = "not-the-start-time" + l.entries[key] = e + l.flushLocked() +} + +// forgetStartTimeForTest makes an entry look like one recorded where the start +// time cannot be read. +func (l *processLedger) forgetStartTimeForTest(key string) { + l.mu.Lock() + defer l.mu.Unlock() + e := l.entries[key] + e.StartTime = "" + l.entries[key] = e + l.flushLocked() +} + +func (l *processLedger) flushLocked() { + list := make([]ledgerEntry, 0, len(l.entries)) + for _, e := range l.entries { + list = append(list, e) + } + data, err := json.Marshal(list) + if err == nil { + err = os.MkdirAll(filepath.Dir(l.path), 0o750) + } + if err == nil { + tmp := l.path + ".tmp" + if err = os.WriteFile(tmp, data, 0o600); err == nil { + err = os.Rename(tmp, l.path) + } + } + if err != nil { + xlog.Warn("Failed to write the worker process ledger", "path", l.path, "error", err) + } +} + +// startTimesReadable reports whether this system can tell a process's start +// time (Linux, through /proc). +func startTimesReadable() bool { return readStartTime(os.Getpid()) != "" } + +// procStartTime returns a process's start time as the kernel reports it, or "" +// when it cannot be read: the process is gone, or there is no /proc (the sweep +// is Linux only; elsewhere an entry is never killed). +var readStartTime = procStartTime + +func procStartTime(pid int) string { + data, err := os.ReadFile("/proc/" + strconv.Itoa(pid) + "/stat") + if err != nil { + return "" + } + // The command name is in parentheses and may hold spaces; the fields after + // the last ')' are space separated and start at field 3. Field 22 is the + // start time. + s := string(data) + i := strings.LastIndex(s, ")") + if i < 0 { + return "" + } + fields := strings.Fields(s[i+1:]) + if len(fields) < 20 { + return "" + } + return fields[19] +} diff --git a/core/services/worker/process_ledger_unix.go b/core/services/worker/process_ledger_unix.go new file mode 100644 index 000000000..92c1c7c0c --- /dev/null +++ b/core/services/worker/process_ledger_unix.go @@ -0,0 +1,8 @@ +//go:build !windows + +package worker + +import "syscall" + +// killProcessGroup sends SIGKILL to the process group led by pid. +func killProcessGroup(pid int) error { return syscall.Kill(-pid, syscall.SIGKILL) } diff --git a/core/services/worker/process_ledger_windows.go b/core/services/worker/process_ledger_windows.go new file mode 100644 index 000000000..97ad03e74 --- /dev/null +++ b/core/services/worker/process_ledger_windows.go @@ -0,0 +1,10 @@ +//go:build windows + +package worker + +import "errors" + +// killProcessGroup is not supported on Windows: backends there are not started +// in process groups, and the ledger never holds a start time (it is read from +// /proc), so no entry is ever swept. +func killProcessGroup(int) error { return errors.New("process groups are not supported on windows") } diff --git a/core/services/worker/registration.go b/core/services/worker/registration.go index e9d93b101..23a439c72 100644 --- a/core/services/worker/registration.go +++ b/core/services/worker/registration.go @@ -247,7 +247,9 @@ func (cfg *Config) registrationBody() map[string]any { // used", while reporting total-as-available lies to the scheduler about // free capacity. func (cfg *Config) heartbeatBody() map[string]any { - body := map[string]any{} + // The incarnation lets the controller learn that this worker restarted, and + // so that every load operation of the previous process ended. + body := map[string]any{"worker_incarnation": workerIncarnation} aggregate := getGPUAggregateInfo() if aggregate.TotalVRAM > 0 { body["available_vram"] = aggregate.FreeVRAM diff --git a/core/services/worker/supervisor.go b/core/services/worker/supervisor.go index 9a5a7c4d3..09c918d4b 100644 --- a/core/services/worker/supervisor.go +++ b/core/services/worker/supervisor.go @@ -5,13 +5,16 @@ import ( "errors" "fmt" "maps" + "net" "os" "path/filepath" "slices" + "strconv" "strings" "sync" "time" + "github.com/google/uuid" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/gallery" "github.com/mudler/LocalAI/core/services/workerctl" @@ -57,6 +60,14 @@ type backendProcess struct { // is exactly what stays the same across a reinstall. backendDir string backendDirID os.FileInfo + + // instance identifies this incarnation of the process. A port can be + // reused by a replacement under the same key, so an address alone does not + // say which process a stop meant. + instance string + // operationID is the load operation this process belongs to, empty when + // none (a legacy start, or a load that completed). + operationID string } const workerBackendFreeTimeout = 5 * time.Second @@ -148,6 +159,17 @@ type backendSupervisor struct { // the same not-yet-cached backend) are serialized here so the gallery // download path doesn't race itself on the same directory. backendLocks map[string]*sync.Mutex + + // operations are the loads the watchdog bounds, by operation id. Guarded by + // mu. See operations.go. + operations map[string]*loadOperation + // ledger records started backends for the orphan sweep. nil in tests. + ledger *processLedger + // opKillTTL, opTick and readyFn are overridden only by tests; zero or nil + // means the defaults. + opKillTTL time.Duration + opTick time.Duration + readyFn func(addr string) bool } // defaultPortQuarantine is how long a released gRPC port waits before it can be @@ -321,6 +343,51 @@ func (s *backendSupervisor) allocatePort(key string) (int, error) { ErrNoFreePort, minPort, maxPort, len(s.processes), len(s.quarantinedPorts)) } +// allocateFreePort is allocatePort that also checks the port is free on the +// host. A restarted worker can be handed a port that an orphan of its +// predecessor still holds. The readiness poll would connect to that orphan and +// report a backend that is not the one it started, so a busy port is set aside +// and the next one is tried. Callers must hold s.mu. +func (s *backendSupervisor) allocateFreePort(key string) (int, error) { + var busy []int + defer func() { + for _, p := range busy { + // Back to the allocator after the quarantine, once nothing holds it. + s.releasePort(p) + } + }() + for range 64 { + port, err := s.allocatePort(key) + if err != nil { + return 0, err + } + if portIsFree(port) { + return port, nil + } + xlog.Warn("A gRPC port is already in use on this host; skipping it", "backend", key, "port", port) + busy = append(busy, port) + } + return 0, fmt.Errorf("%w: every port tried was already in use", ErrNoFreePort) +} + +// portIsFree reports whether nothing listens on port on this host. Two checks, +// because neither is enough everywhere. A connect to loopback finds a listener +// bound to a specific address, which a bind of the wildcard address can miss on +// BSD and macOS (the listener sets SO_REUSEADDR). A bind of the wildcard address +// finds one that does not answer a connect. +func portIsFree(port int) bool { + if conn, err := net.DialTimeout("tcp", fmt.Sprintf("127.0.0.1:%d", port), 200*time.Millisecond); err == nil { + _ = conn.Close() + return false + } + lis, err := net.Listen("tcp", fmt.Sprintf("0.0.0.0:%d", port)) + if err != nil { + return false + } + _ = lis.Close() + return true +} + // sweepAffinity drops claims whose window has lapsed, so their ports become // ordinary free ports again. Swept lazily on allocation for the same reason as // sweepQuarantine: the only observer is allocation itself, so a timer goroutine @@ -462,7 +529,7 @@ func (s *backendSupervisor) startBackend(backend, backendName, backendPath strin s.reapDeadProcess(backend, bp) } - port, err := s.allocatePort(backend) + port, err := s.allocateFreePort(backend) if err != nil { s.mu.Unlock() return "", fmt.Errorf("allocating gRPC port for backend %s: %w", backend, err) @@ -496,6 +563,10 @@ func (s *backendSupervisor) startBackend(backend, backendName, backendPath strin backendName: backendName, backendDir: backendDir, backendDirID: dirInfo, + instance: uuid.NewString(), + } + if pid, convErr := strconv.Atoi(proc.CurrentPID()); convErr == nil { + s.ledger.add(backend, pid) } xlog.Info("Backend process started", "backend", backend, "addr", clientAddr) @@ -521,6 +592,12 @@ func (s *backendSupervisor) startBackend(backend, backendName, backendPath strin time.Sleep(readinessPollInterval) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) ok, healthErr := client.HealthCheck(ctx) + // An answer only counts when the process this worker started is still + // alive. A child that lost the bind to an orphan exits at once, and the + // orphan's answer is not its own. + if ok && !proc.IsAlive() { + ok = false + } if ok { cancel() // Verify the process wasn't stopped/replaced while health-checking. @@ -599,6 +676,7 @@ func (s *backendSupervisor) markBackendServing(key string, bp *backendProcess) b func (s *backendSupervisor) reapDeadProcess(key string, bp *backendProcess) { xlog.Warn("Backend process died unexpectedly, restarting", "backend", key) delete(s.processes, key) + s.ledger.remove(key) if bp == nil { return } @@ -620,6 +698,7 @@ func (s *backendSupervisor) releaseBackendStart(key string, bp *backendProcess) return } delete(s.processes, key) + s.ledger.remove(key) s.cleanupProcessRuntime(bp.proc) if bp.port <= 0 { xlog.Error("Cannot recycle backend port: startup has invalid recorded port", "backend", key, "addr", bp.addr, "port", bp.port) @@ -850,7 +929,25 @@ func (s *backendSupervisor) stopBackendExact(key string, force bool) error { if bp == nil { return nil } + return s.finishStopping(key, bp, force) +} +// stopBackendExactBP stops exactly bp, and only if it is still the process the +// supervisor holds under key. The watchdog uses it: it chose its victim earlier +// and the key may have been reused since. +func (s *backendSupervisor) stopBackendExactBP(key string, bp *backendProcess, force bool) error { + s.mu.Lock() + current, ok := s.processes[key] + if !ok || current != bp || bp.proc == nil || bp.stopping { + s.mu.Unlock() + return nil + } + bp.stopping = true + s.mu.Unlock() + return s.finishStopping(key, bp, force) +} + +func (s *backendSupervisor) finishStopping(key string, bp *backendProcess, force bool) error { if !force { client := grpc.NewClientWithToken(bp.addr, false, nil, false, s.cfg.RegistrationToken) freeCtx, cancel := context.WithTimeout(context.Background(), workerBackendFreeTimeout) @@ -862,6 +959,7 @@ func (s *backendSupervisor) stopBackendExact(key string, force bool) error { } xlog.Info("Stopping backend process", "backend", key, "addr", bp.addr, "force", force, "backendName", bp.backendName) + s.ml.NoteIntentionalStop(bp.proc) stopErr := bp.proc.Stop() if stopErr != nil { xlog.Error("Error stopping backend process", "backend", key, "error", stopErr) @@ -878,13 +976,20 @@ func (s *backendSupervisor) stopModelExact(req workerctl.ModelStopRequest) worke s.mu.Lock() bp, ok := s.processes[req.ProcessKey] if !ok || bp.proc == nil { + if op, known := s.operations[req.OperationID]; known && req.OperationID != "" { + op.expired = true + delete(s.operations, req.OperationID) + } s.mu.Unlock() reply.Terminated = true return reply } reply.Matched = true reply.Address = bp.addr - if bp.addr != req.ExpectedAddress { + // A stop of a load operation may omit the address: the controller learns it + // only after the install replies. The operation, process key and instance + // checks below then carry the identity. + if bp.addr != req.ExpectedAddress && !(req.OperationID != "" && req.ExpectedAddress == "") { s.mu.Unlock() reply.Error = fmt.Sprintf("address mismatch for process %s: recorded %q, expected %q", req.ProcessKey, bp.addr, req.ExpectedAddress) return reply @@ -894,6 +999,17 @@ func (s *backendSupervisor) stopModelExact(req workerctl.ModelStopRequest) worke reply.Error = fmt.Sprintf("process %s is already stopping", req.ProcessKey) return reply } + if req.OperationID != "" { + if err := s.checkOperationTarget(req, bp); err != nil { + s.mu.Unlock() + reply.Error = err.Error() + return reply + } + } else if req.ProcessInstance != "" && bp.instance != req.ProcessInstance { + s.mu.Unlock() + reply.Error = fmt.Sprintf("process instance mismatch for %s", req.ProcessKey) + return reply + } bp.stopping = true s.mu.Unlock() @@ -909,6 +1025,7 @@ func (s *backendSupervisor) stopModelExact(req workerctl.ModelStopRequest) worke } } + s.ml.NoteIntentionalStop(bp.proc) stopErr := bp.proc.Stop() if stopErr == nil { <-bp.proc.Done() @@ -921,6 +1038,14 @@ func (s *backendSupervisor) stopModelExact(req workerctl.ModelStopRequest) worke } return reply } + if req.OperationID != "" { + s.mu.Lock() + if op, known := s.operations[req.OperationID]; known { + op.expired = true + delete(s.operations, req.OperationID) + } + s.mu.Unlock() + } reply.Terminated = true return reply } @@ -954,6 +1079,7 @@ func (s *backendSupervisor) finishBackendStop(key string, bp *backendProcess, st return fmt.Errorf("stopping backend process %s: %w", key, stopErr) } delete(s.processes, key) + s.ledger.remove(key) s.cleanupProcessRuntime(bp.proc) if bp.port <= 0 { xlog.Error("Cannot recycle backend port: process has invalid recorded port", "backend", key, "addr", bp.addr, "port", bp.port) diff --git a/core/services/worker/worker.go b/core/services/worker/worker.go index 60d469574..79df4e884 100644 --- a/core/services/worker/worker.go +++ b/core/services/worker/worker.go @@ -257,6 +257,16 @@ func Run(ctx *cliContext.Context, cfg *Config) error { }), )) + // A previous worker that was killed left its backends running. Kill them + // now, before this worker hands out ports or reports an incarnation. + supervisor.ledger = newProcessLedger(filepath.Join(dataDir, "worker-processes.json")) + if n := supervisor.ledger.sweepStale(); n > 0 { + xlog.Warn("Killed backend process groups left behind by a previous worker", "count", n) + } + + // The watchdog stops load operations the controller no longer renews. + go supervisor.runOperationWatchdog(shutdownCtx) + control := newNATSControlServer(natsClient, nodeID) if err := supervisor.registerLifecycleVerbs(control); err != nil { nodes.ShutdownFileTransferServer(httpServer) diff --git a/core/services/workerctl/backend.go b/core/services/workerctl/backend.go index a11c3fb62..6321b12ae 100644 --- a/core/services/workerctl/backend.go +++ b/core/services/workerctl/backend.go @@ -29,6 +29,17 @@ type BackendInstallRequest struct { // running, debounced to roughly 250ms. Empty means the caller is a // reconciler-driven retry that does not need progress streamed. OpID string `json:"op_id,omitempty"` + // OperationID names the load this install belongs to (the load job + // generation). With it the worker tracks the load as an operation it can + // bound: it kills the backend when renewals stop or the deadline passes. + // Workers older than this field ignore it. Empty means a controller older + // than this field: the worker then tracks an anonymous operation. + OperationID string `json:"operation_id,omitempty"` + // DeadlineMs is the longest the load may run, as a duration in + // milliseconds, not a timestamp, so worker clock skew does not matter. The + // worker converts it to its own monotonic deadline when the request + // arrives. Zero means the worker's default. + DeadlineMs int64 `json:"deadline_ms,omitempty"` } // BackendInstallReply is the response from a backend.install control request. @@ -36,6 +47,15 @@ type BackendInstallReply struct { Success bool `json:"success"` Address string `json:"address,omitempty"` // gRPC address of the backend process (host:port) Error string `json:"error,omitempty"` + // ProcessInstance identifies this incarnation of the backend process, so a + // later stop cannot hit a replacement that took the same port. + ProcessInstance string `json:"process_instance,omitempty"` + // ReportsOperations is true on a worker that tracks load operations. A reply + // without it comes from a worker that predates them: it cannot confirm a + // stop, so the controller stops its process by exact address and holds the + // model for the load deadline. The capability belongs to the worker, not to + // the transport. + ReportsOperations bool `json:"reports_operations,omitempty"` } // BackendUpgradeRequest is the payload for a backend.upgrade control request. diff --git a/core/services/workerctl/model.go b/core/services/workerctl/model.go index b9074bb63..2ae450883 100644 --- a/core/services/workerctl/model.go +++ b/core/services/workerctl/model.go @@ -6,6 +6,12 @@ type ModelStopRequest struct { ExpectedAddress string `json:"expected_address"` Force bool `json:"force,omitempty"` ConfigRevision string `json:"config_revision,omitempty"` + // OperationID and ProcessInstance address one load operation. When set, the + // worker checks them, together with ProcessKey and ExpectedAddress, against + // its own records and refuses on any mismatch. The worker never falls back + // to a name-only match. + OperationID string `json:"operation_id,omitempty"` + ProcessInstance string `json:"process_instance,omitempty"` } type ModelStopReply struct { @@ -21,6 +27,9 @@ type ModelStopReply struct { type ModelUnloadRequest struct { ModelName string `json:"model_name"` Address string `json:"address,omitempty"` // gRPC address of the backend process to unload from + // ProcessInstance, when set, must match the process at Address. The worker + // never picks a process on its own: a request with no Address unloads nothing. + ProcessInstance string `json:"process_instance,omitempty"` } // ModelUnloadReply is the response from a model.unload control request. @@ -47,6 +56,9 @@ type ModelsRunningRequest struct{} type ModelsRunningReply struct { Models []RunningModelInfo `json:"models"` Error string `json:"error,omitempty"` + // ReportsOperations is true on workers that track load operations. Without + // it a controller must treat the node as legacy: it cannot confirm a stop. + ReportsOperations bool `json:"reports_operations,omitempty"` } // RunningModelInfo identifies one live backend process on a worker. The triple @@ -56,4 +68,28 @@ type RunningModelInfo struct { ModelID string `json:"model_id"` ReplicaIndex int `json:"replica_index"` Address string `json:"address,omitempty"` + // OperationID is the load operation that started this process, empty for a + // process with none (a legacy start). ProcessInstance identifies this + // incarnation of the process. + OperationID string `json:"operation_id,omitempty"` + ProcessInstance string `json:"process_instance,omitempty"` +} + +// OperationRequest renews and completes load operations on a worker. The +// controller sends it on the same tick as its own lease heartbeat, batched per +// node. A missed renewal costs nothing until the worker's kill TTL. +type OperationRequest struct { + // Renew extends each operation's lease on the worker. + Renew []string `json:"renew,omitempty"` + // Complete ends each operation: the load finished and the backend now + // serves, so the watchdog must stop watching it. + Complete []string `json:"complete,omitempty"` +} + +// OperationReply names the operations the worker does not know, so the +// controller can tell a lost operation from a renewed one. +type OperationReply struct { + Renewed []string `json:"renewed,omitempty"` + Completed []string `json:"completed,omitempty"` + Unknown []string `json:"unknown,omitempty"` } diff --git a/docs/content/features/distributed-mode.md b/docs/content/features/distributed-mode.md index 3d99cd63b..2fe4d4f7f 100644 --- a/docs/content/features/distributed-mode.md +++ b/docs/content/features/distributed-mode.md @@ -124,6 +124,11 @@ So the load does **not** run on the request. The first request for an unloaded m - It never starts a duplicate load and never blocks on the database lock. (Before this split, concurrent requests blocked on `pg_advisory_lock` for the whole load and were killed by the PostgreSQL role's `statement_timeout` — `SQLSTATE 57014` — so from the operator's seat the model simply never loaded.) - If the load fails, the waiter gets the *real* cause (`worker out of disk`), not an anonymous timeout. - If the client disconnects, the load keeps going. It belongs to the job record, not to the request. +- The owner of a job holds a 30 second lease, which it renews on every heartbeat. The database clock decides whether a lease has expired, so a frontend with a wrong clock cannot expire a live lease or keep a dead one. If an owner cannot renew for a whole lease, it stops its own load. +- If an owner dies, its job is marked failed once the lease runs out, with or without a new request. A request that arrives then reads the cause. The model is held for a 2.5 minute stop window, because remote work may still run. After that the next request starts a new attempt. No manual cleanup is needed. +- A failure that is known to have ended the remote work (an error from the backend, or a failure before a node was chosen) is kept for 15 seconds only, so every waiter reads the same cause and the next request can retry. +- Each attempt has its own generation. If a job is replaced, the old owner notices at its next heartbeat and stops its load. Its late writes to the job and to the replica table are rejected. +- If the job table cannot be read, a cold load fails instead of running without a job. Models that are already loaded keep serving, because routing to a loaded replica does not read the job table. When the wait budget (`LOCALAI_MODEL_LOAD_WAIT`, default `60s`) runs out, the request is answered with `503`, a `Retry-After` header, and a body that says exactly where the load is: @@ -153,9 +158,54 @@ The `error` envelope keeps OpenAI clients working unchanged; `loading` is additi The chat UI renders this state inline and retries automatically once the model reports ready. Poll `GET /api/models/{id}/load-status` for the same `loading` object at any time. {{% notice note %}} -A frontend replica that dies mid-load does not wedge the model: the job row carries a heartbeat and another replica reclaims a job whose heartbeat has stopped. The heartbeat is time-based, not byte-based, because a checkpoint load legitimately transfers zero bytes for many minutes. +A frontend replica that dies mid-load does not wedge the model: the job row carries a lease, and a job whose lease ran out is failed and then released. The lease is renewed on a timer, not on byte progress, because a checkpoint load legitimately transfers zero bytes for many minutes. {{% /notice %}} +#### The worker bounds the work it runs + +The frontend owns the job row, but the real work runs in a backend process on a worker. The worker therefore watches each load too. A load is an **operation** named by the job's generation: + +- The install request carries the operation id and the longest the load may run, as a duration, so a worker clock that is wrong changes nothing. The backend starts in its own process group. +- The frontend renews the operation every few seconds and completes it when the load finishes. The worker kills the whole process group when no renewal arrives for 90 seconds, or when the deadline passes. A backend that already reports `READY` is never killed: a lost completion message must not destroy a model that serves. +- A stop names the operation, the process key and, when known, the address and process instance. The worker refuses unless they all match its own records. There is no fallback to "any running backend". A stop for a load that already finished leaves the serving model alone. +- A worker that is killed cannot stop its backends. It records each backend's process group in a small file under its data directory, and the next worker kills every group listed there before it serves (Linux; the start time of the leader guards against a recycled pid). A new incarnation, reported on the next heartbeat, then confirms the failed loads on that node. The worker also skips a gRPC port that something already listens on, and a readiness answer only counts while the worker's own backend process is alive. +- If the worker answers a renewal with "unknown operation" three times in a row, the frontend fails the load at once instead of waiting for the load budget. The worker lost the operation, so the work is gone. + +When a load fails after remote work may have started (a timeout, a cancel, a lost lease), the owner stops the operation immediately. An acknowledged stop shortens the hold to the 15 second report window. A silent worker keeps the hold at the stop window (2.5 minutes), and the reconciler retries the stop on every pass until the worker answers or the window ends. Nothing needs manual cleanup. + +| Setting | Value | Meaning | +|---------|-------|---------| +| Lease TTL | 30 s | How long a job's lease lasts after each renewal | +| Worker kill TTL | 90 s | No renewal for this long: the worker kills the operation | +| Stop window | 150 s | How long a failed load holds the model if the worker never confirms | +| Report window | 15 s | How long a failure with confirmed-ended work is kept | + +A model held by a failed job answers `503` with `Retry-After` set to the seconds until the hold ends, and the real cause in the body. + +#### Cancelling a load + +`POST /api/models/{id}/load-cancel` (admin only) cancels one load attempt. The body names the exact attempt, as `GET /api/models/{id}/load-status` reports it: + +```json +{"job_id": "0b6e4a3c-5c1d-4d52-8f0a-0f3c9e0b8f11"} +``` + +| Status | Meaning | +|--------|---------| +| `200` | `state: stopped` (the worker confirmed) or `state: gone` (no such load any more) | +| `202` | `state: stopping`. The cancel is recorded and the stop is pending. The model is released after `retry_after` seconds regardless. | +| `400` | The body is not `{"job_id": "..."}` | +| `404` | Unknown model, or the server is not distributed | +| `409` | A different attempt is current. The body carries its `current_job_id`. | + +The call is idempotent. A repeat retries the stop and never extends the hold. A load that has not been placed on a node yet can be cancelled too. Unloading a model on a node, draining a node, and removing a node all cancel the loads placed there through the same stop path, and an unload still unloads the loaded replicas. The replica rows of a cancelled attempt are removed as soon as the worker confirms the stop. The `cancel_model_load` tool of the assistant calls the same service. + +`load-status` also reports `job_id`, `lease_expires_in`, `cancel_requested`, `last_error`, `stopping`, `stop_deadline` and `retry_after`. A database error is a `503`, never an empty answer. + +#### Rolling upgrades + +Upgrade the frontends first. A worker that predates operations ignores the new request fields and does not report `reports_operations`. The frontend then treats the node as legacy: it cannot confirm a stop, so a failed load holds the model for the 45 minute load deadline, as it did before leases existed, and never longer. For such a node the stop, including a cancel, is sent by exact process address, never by model name. If the address is not known, no stop is claimed, and the model is held for the 45 minutes. A new worker that gets an install from an older frontend tracks it as an anonymous operation: it kills it at its deadline only, never for missing renewals. + ### NATS JWT authentication (recommended for production) By default, NATS connections are anonymous: any client that can reach port `4222` may publish control-plane subjects such as `nodes..backend.install`. Enable JWT auth to scope workers to their own node subjects and give the frontend a dedicated service credential. @@ -544,7 +594,8 @@ Used by the WebUI and admin API consumers. Requires admin authentication. | `POST` | `/api/nodes/:id/backends/install` | Install a backend on a worker | | `POST` | `/api/nodes/:id/backends/upgrade` | Upgrade (force-reinstall) a backend on a worker | | `POST` | `/api/nodes/:id/backends/delete` | Delete a backend from a worker | -| `POST` | `/api/nodes/:id/models/unload` | Unload a model from a worker | +| `POST` | `/api/nodes/:id/models/unload` | Unload a model from a worker. Cancels a load of that model on the worker first. | +| `POST` | `/api/models/:id/load-cancel` | Cancel one load attempt (`{"job_id": "..."}`) | | `POST` | `/api/nodes/:id/models/delete` | Delete model files from a worker | | `PUT` | `/api/nodes/:id/vram-budget` | Set a VRAM budget for a worker (`{"value":"80%"}`) | | `DELETE` | `/api/nodes/:id/vram-budget` | Clear a worker's VRAM budget (revert to all detected VRAM) | @@ -1314,7 +1365,7 @@ Notes: **A model cannot be scheduled on a node that looks free (`no replica slot ... all models busy, cannot evict`):** - A replica row in `staging` or `loading` holds its slot: slot allocation counts every state except `unloading`. If a worker drops out mid-transfer, that row never reaches `loaded`, and eviction only ever considers `loaded` replicas, so on a node with one replica slot per model the model became unschedulable there. - The reconciler now reclaims a replica row stuck before serving when no load job is still driving it, and the freed slot is immediately reusable. -- Liveness is decided by the load job's progress heartbeat, not by elapsed time. Staging a large checkpoint legitimately runs for a long time without touching the replica row, so a transfer that is still progressing is never reclaimed however long it takes. +- Each replica row names the load attempt that made it. Liveness is decided by that attempt's job lease, not by elapsed time. Staging a large checkpoint legitimately runs for a long time without touching the replica row, so a transfer whose owner still renews its lease is never reclaimed however long it takes. A row is reclaimed when its attempt has no job: the job was released after a failure, or another attempt replaced it. - `Reconciler: reclaimed a replica slot held by a load nobody is driving` names each row reclaimed this way. **A request fails with `nats: no responders available for request`:** diff --git a/pkg/mcp/localaitools/client.go b/pkg/mcp/localaitools/client.go index 7cba8a8de..ba5c2834a 100644 --- a/pkg/mcp/localaitools/client.go +++ b/pkg/mcp/localaitools/client.go @@ -40,6 +40,9 @@ type LocalAIClient interface { // it down). For a realtime pipeline model every configured sub-model is // loaded; it returns the model names that became resident. LoadModel(ctx context.Context, model string) ([]string, error) + // CancelModelLoad cancels one distributed load attempt, named by the job id + // that load-status reports. It never cancels a replacement attempt. + CancelModelLoad(ctx context.Context, model, jobID string) (LoadCancelResult, error) ImportModelURI(ctx context.Context, req ImportModelURIRequest) (*ImportModelURIResponse, error) // ---- Model aliases ---- diff --git a/pkg/mcp/localaitools/coverage_test.go b/pkg/mcp/localaitools/coverage_test.go index 75831769e..6ec572024 100644 --- a/pkg/mcp/localaitools/coverage_test.go +++ b/pkg/mcp/localaitools/coverage_test.go @@ -54,6 +54,7 @@ var toolToHTTPRoute = map[string]string{ ToolDeleteModel: "POST /models/delete/:name", ToolEditModelConfig: "PATCH /api/models/config-json/:name", ToolReloadModels: "POST /models/reload", + ToolCancelModelLoad: "POST /api/models/:id/load-cancel", ToolLoadModel: "POST /backend/load", ToolInstallBackend: "POST /backends/apply", ToolUpgradeBackend: "POST /backends/upgrade/:name", diff --git a/pkg/mcp/localaitools/dto.go b/pkg/mcp/localaitools/dto.go index afd19bafd..88d7159ff 100644 --- a/pkg/mcp/localaitools/dto.go +++ b/pkg/mcp/localaitools/dto.go @@ -436,3 +436,13 @@ type FailoverChainInfo struct { Pinned string `json:"pinned,omitempty"` Targets []FailoverTargetInfo `json:"targets"` } + +// LoadCancelResult is what a cancel of a distributed load did. State is +// "stopped" (the worker confirmed), "stopping" (recorded, stop pending; the model +// is released after RetryAfter seconds regardless) or "gone". +type LoadCancelResult struct { + Model string `json:"model"` + JobID string `json:"job_id"` + State string `json:"state"` + RetryAfter int `json:"retry_after,omitempty"` +} diff --git a/pkg/mcp/localaitools/fakes_test.go b/pkg/mcp/localaitools/fakes_test.go index cba6efa57..d8a85491d 100644 --- a/pkg/mcp/localaitools/fakes_test.go +++ b/pkg/mcp/localaitools/fakes_test.go @@ -69,6 +69,11 @@ type fakeCall struct { args any } +func (f *fakeClient) CancelModelLoad(_ context.Context, model, jobID string) (LoadCancelResult, error) { + f.record("CancelModelLoad", []string{model, jobID}) + return LoadCancelResult{Model: model, JobID: jobID, State: "stopping"}, nil +} + func (f *fakeClient) record(method string, args any) { f.mu.Lock() defer f.mu.Unlock() diff --git a/pkg/mcp/localaitools/httpapi/client.go b/pkg/mcp/localaitools/httpapi/client.go index 631977c03..726f121a3 100644 --- a/pkg/mcp/localaitools/httpapi/client.go +++ b/pkg/mcp/localaitools/httpapi/client.go @@ -846,6 +846,12 @@ func (c *Client) PinFailoverTarget(ctx context.Context, chain, target string) er return c.do(ctx, http.MethodPost, routeFailover+"/"+url.PathEscape(chain)+"/pin", map[string]string{"target": target}, nil) } +func (c *Client) CancelModelLoad(ctx context.Context, model, jobID string) (localaitools.LoadCancelResult, error) { + var result localaitools.LoadCancelResult + err := c.do(ctx, http.MethodPost, "/api/models/"+url.PathEscape(model)+"/load-cancel", map[string]string{"job_id": jobID}, &result) + return result, err +} + func (c *Client) UnpinFailoverTarget(ctx context.Context, chain string) error { return c.do(ctx, http.MethodDelete, routeFailover+"/"+url.PathEscape(chain)+"/pin", nil, nil) } diff --git a/pkg/mcp/localaitools/inproc/client.go b/pkg/mcp/localaitools/inproc/client.go index ff80b6cac..b1900656b 100644 --- a/pkg/mcp/localaitools/inproc/client.go +++ b/pkg/mcp/localaitools/inproc/client.go @@ -46,12 +46,15 @@ import ( // distributed-aware, ModelConfigLoader manages on-disk YAML, etc.), so this // layer just translates between MCP DTOs and service signatures. type Client struct { - AppConfig *config.ApplicationConfig - SystemState *system.SystemState - ConfigLoader *config.ModelConfigLoader - ModelLoader *model.ModelLoader - Gallery *galleryop.GalleryService - NodeRegistry *nodes.NodeRegistry + AppConfig *config.ApplicationConfig + SystemState *system.SystemState + ConfigLoader *config.ModelConfigLoader + ModelLoader *model.ModelLoader + Gallery *galleryop.GalleryService + NodeRegistry *nodes.NodeRegistry + // LoadStopper stops the remote work of a cancelled load. It is the same + // stopper the HTTP endpoint uses, so both paths stop work the same way. + LoadStopper nodes.LoadAttemptStopper VoiceProfiles *voiceprofile.Store // StatsRecorder and FallbackUser are optional — they back the @@ -1146,3 +1149,26 @@ func (c *Client) UnpinFailoverTarget(_ context.Context, chain string) error { } return c.Failover.Unpin(chain) } + +// CancelModelLoad cancels one distributed load attempt through the same service +// the load-cancel endpoint uses, with the same stopper. +func (c *Client) CancelModelLoad(ctx context.Context, model, jobID string) (localaitools.LoadCancelResult, error) { + result := localaitools.LoadCancelResult{Model: model, JobID: jobID} + if model == "" || jobID == "" { + return result, errors.New("model and job_id are required") + } + if c.NodeRegistry == nil { + return result, errors.New("load cancellation is only available in distributed mode") + } + svc := &nodes.LoadCancelService{Registry: c.NodeRegistry, Stopper: c.LoadStopper} + out, err := svc.Cancel(ctx, nodes.LoadJobRef{TrackingKey: model, Generation: jobID}) + if errors.Is(err, nodes.ErrLoadCancelConflict) { + return result, fmt.Errorf("a different load attempt is current (job_id %s); read load-status again", out.CurrentJobID) + } + if err != nil { + return result, err + } + result.State = string(out.State) + result.RetryAfter = int(out.RetryAfter.Seconds()) + return result, nil +} diff --git a/pkg/mcp/localaitools/prompts/10_safety.md b/pkg/mcp/localaitools/prompts/10_safety.md index 5f9c5f840..1417a51cf 100644 --- a/pkg/mcp/localaitools/prompts/10_safety.md +++ b/pkg/mcp/localaitools/prompts/10_safety.md @@ -2,7 +2,7 @@ These rules are non-negotiable. The user trusts you to operate their server without unintended changes. -1. **Confirm before mutating.** Before calling any of these tools — `install_model`, `import_model_uri`, `delete_model`, `install_backend`, `upgrade_backend`, `edit_model_config`, `reload_models`, `load_model`, `toggle_model_state`, `toggle_model_pinned`, `set_branding`, `set_alias`, `seed_router_corpus`, `clear_router_corpus`, `create_voice_profile`, `delete_voice_profile`, `set_node_vram_budget`, `set_scheduling`, `delete_scheduling`, `pin_failover_target`, `unpin_failover_target` — first state in plain language what you are about to do (which tool, which target, which arguments) and wait for the user's explicit confirmation in the next turn. "Yes", "do it", "go ahead", "proceed" all count as confirmation. Anything else does not. +1. **Confirm before mutating.** Before calling any of these tools — `install_model`, `import_model_uri`, `delete_model`, `install_backend`, `upgrade_backend`, `edit_model_config`, `reload_models`, `load_model`, `cancel_model_load`, `toggle_model_state`, `toggle_model_pinned`, `set_branding`, `set_alias`, `seed_router_corpus`, `clear_router_corpus`, `create_voice_profile`, `delete_voice_profile`, `set_node_vram_budget`, `set_scheduling`, `delete_scheduling`, `pin_failover_target`, `unpin_failover_target` — first state in plain language what you are about to do (which tool, which target, which arguments) and wait for the user's explicit confirmation in the next turn. "Yes", "do it", "go ahead", "proceed" all count as confirmation. Anything else does not. 2. **Disambiguate before mutating.** If the user's request is ambiguous (several gallery candidates match, the model name has multiple installed versions, the backend has variants), present the candidates as a numbered list and ask the user to pick before calling any mutating tool. diff --git a/pkg/mcp/localaitools/prompts/20_tools.md b/pkg/mcp/localaitools/prompts/20_tools.md index 88404801e..591c49f51 100644 --- a/pkg/mcp/localaitools/prompts/20_tools.md +++ b/pkg/mcp/localaitools/prompts/20_tools.md @@ -35,6 +35,7 @@ The MCP `tools/list` endpoint also exposes the full input schema for each of the - `upgrade_backend` — Upgrade an installed backend by name. - `edit_model_config` — Patch (deep-merge) JSON into an installed model's config. - `reload_models` — Reload all model configs from disk. +- `cancel_model_load` — Cancel one distributed load by the exact job_id from load-status. `stopping` means pending, not proof the work stopped. - `load_model` — Pre-load a model into memory so the first request pays no cold-start cost. For a realtime pipeline model, every sub-model (VAD, transcription, LLM, TTS, sound_detection, voice_recognition) is loaded. Inverse of stopping a model. - `toggle_model_state` — Enable or disable a model (`action`: `enable` or `disable`). - `toggle_model_pinned` — Pin or unpin a model (`action`: `pin` or `unpin`). diff --git a/pkg/mcp/localaitools/prompts/skills/cancel_model_load.md b/pkg/mcp/localaitools/prompts/skills/cancel_model_load.md new file mode 100644 index 000000000..c6ca0089b --- /dev/null +++ b/pkg/mcp/localaitools/prompts/skills/cancel_model_load.md @@ -0,0 +1,7 @@ +# Skill: Cancel a Distributed Load + +1. Read the model's load-status (`GET /api/models/{id}/load-status`) and take the exact `job_id` from it. Never guess a job id. +2. Say what you will cancel and ask for confirmation under safety rule 1. +3. Call `cancel_model_load` with that model and job id, only after the confirmation. +4. Read the state. `stopped` means the worker confirmed the work ended. `stopping` means the cancel is recorded and the stop is pending; the model is free again after `retry_after` seconds regardless. `gone` means no such load exists any more. +5. On a conflict, a different attempt is now current. Read load-status again and ask before cancelling that one. Do not treat a conflict or an unknown model as success. diff --git a/pkg/mcp/localaitools/server_test.go b/pkg/mcp/localaitools/server_test.go index 6fe104b9e..cc33024ec 100644 --- a/pkg/mcp/localaitools/server_test.go +++ b/pkg/mcp/localaitools/server_test.go @@ -149,6 +149,7 @@ var _ = Describe("Tool dispatch", func() { {ToolUpgradeBackend, map[string]any{"name": "llama-cpp"}, "UpgradeBackend"}, {ToolEditModelConfig, map[string]any{"name": "foo", "patch": map[string]any{"context_size": 4096}}, "EditModelConfig"}, {ToolReloadModels, struct{}{}, "ReloadModels"}, + {ToolCancelModelLoad, map[string]any{"model": "test-model", "job_id": "generation"}, "CancelModelLoad"}, {ToolLoadModel, map[string]any{"model": "test-model"}, "LoadModel"}, {ToolToggleModelState, map[string]any{"name": "foo", "action": "enable"}, "ToggleModelState"}, {ToolToggleModelPinned, map[string]any{"name": "foo", "action": "pin"}, "ToggleModelPinned"}, diff --git a/pkg/mcp/localaitools/tools.go b/pkg/mcp/localaitools/tools.go index e2c9b32d3..ee03d196e 100644 --- a/pkg/mcp/localaitools/tools.go +++ b/pkg/mcp/localaitools/tools.go @@ -35,6 +35,7 @@ const ( ToolDeleteModel = "delete_model" ToolEditModelConfig = "edit_model_config" ToolReloadModels = "reload_models" + ToolCancelModelLoad = "cancel_model_load" ToolLoadModel = "load_model" ToolInstallBackend = "install_backend" ToolUpgradeBackend = "upgrade_backend" @@ -78,6 +79,7 @@ var mutatingToolNames = []string{ ToolDeleteModel, ToolEditModelConfig, ToolReloadModels, + ToolCancelModelLoad, ToolLoadModel, ToolInstallBackend, ToolUpgradeBackend, diff --git a/pkg/mcp/localaitools/tools_config.go b/pkg/mcp/localaitools/tools_config.go index 0f8663d9a..3ab9cc94b 100644 --- a/pkg/mcp/localaitools/tools_config.go +++ b/pkg/mcp/localaitools/tools_config.go @@ -41,6 +41,20 @@ func registerConfigTools(s *mcp.Server, client LocalAIClient, opts Options) { return } + mcp.AddTool(s, &mcp.Tool{Name: ToolCancelModelLoad, Description: "Cancel one distributed model load, named by the job_id from the model's load-status. Requires user confirmation per safety rule 1. State `stopping` means the cancel is recorded and the stop is pending, not that the work stopped; the model is released after retry_after seconds regardless."}, func(ctx context.Context, _ *mcp.CallToolRequest, args struct { + Model string `json:"model" jsonschema:"The model whose load to cancel."` + JobID string `json:"job_id" jsonschema:"The exact load attempt, from load-status. Never guess it."` + }) (*mcp.CallToolResult, any, error) { + if args.Model == "" || args.JobID == "" { + return errorResultf("model and job_id are required"), nil, nil + } + result, err := client.CancelModelLoad(ctx, args.Model, args.JobID) + if err != nil { + return errorResult(err), nil, nil + } + return jsonResult(result), nil, nil + }) + mcp.AddTool(s, &mcp.Tool{ Name: ToolEditModelConfig, Description: "Patch (deep-merge) JSON into an installed model's config. Requires user confirmation per safety rule 1; show a diff first.", diff --git a/pkg/model/initializers.go b/pkg/model/initializers.go index 2ae9242e3..cf5003912 100644 --- a/pkg/model/initializers.go +++ b/pkg/model/initializers.go @@ -352,7 +352,14 @@ func (ml *ModelLoader) backendLoader(opts ...Option) (client grpc.Backend, err e if stopErr := ml.StopGRPC(only(o.modelID)); stopErr != nil { xlog.Debug("cleanup stop after failed load", "error", stopErr, "model", o.modelID) } - xlog.Error("Failed to load model", "modelID", o.modelID, "error", err, "backend", o.backendString) + // A model held by a failed distributed load, or still loading, is an + // answer to retry later, not a fault of this request. + var retry interface{ RetryLater() bool } + if errors.As(err, &retry) && retry.RetryLater() { + xlog.Info("Model is not available yet", "modelID", o.modelID, "reason", err) + } else { + xlog.Error("Failed to load model", "modelID", o.modelID, "error", err, "backend", o.backendString) + } return nil, err } diff --git a/pkg/model/process.go b/pkg/model/process.go index ee65de06e..ef67fabc2 100644 --- a/pkg/model/process.go +++ b/pkg/model/process.go @@ -407,6 +407,15 @@ func (ml *ModelLoader) startProcess(grpcProcess, id string, serverAddress string return grpcControlProcess, nil } +// NoteIntentionalStop marks a process that its owner is about to stop on +// purpose, so its exit is logged as a stop and not as a crash. Callers that stop +// a process directly, without going through the loader, use it. +func (ml *ModelLoader) NoteIntentionalStop(proc *process.Process) { + if ml != nil && proc != nil { + ml.stoppingProcs.Store(proc, struct{}{}) + } +} + // forgetExitedProcess drops the model store entry of a backend that exited on // its own, so it is no longer reported as loaded. It shares the lifecycle lock // with loading and shutdown and matches the process identity, so a late exit diff --git a/swagger/docs.go b/swagger/docs.go index 13c9304f7..759b0fdc9 100644 --- a/swagger/docs.go +++ b/swagger/docs.go @@ -1278,9 +1278,86 @@ const docTemplate = `{ } } }, + "/api/models/{id}/load-cancel": { + "post": { + "description": "Cancels the load attempt named by ` + "`" + `job_id` + "`" + ` and stops its remote work. 200 means the attempt is gone or the worker confirmed the stop. 202 means the cancel is recorded and the stop is pending; the model is released after ` + "`" + `retry_after` + "`" + ` seconds regardless. 409 means a different attempt is current and carries its ` + "`" + `current_job_id` + "`" + `. Repeating the call is safe and does not extend the hold.", + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "models" + ], + "summary": "Cancel one distributed model load.", + "parameters": [ + { + "type": "string", + "description": "Model ID", + "name": "id", + "in": "path", + "required": true + }, + { + "description": "The exact attempt to cancel", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelRequest" + } + } + ], + "responses": { + "200": { + "description": "Stopped, or no such load any more", + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelResponse" + } + }, + "202": { + "description": "Cancel recorded; stop pending", + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelResponse" + } + }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "403": { + "description": "Forbidden", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "404": { + "description": "Unknown model", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "409": { + "description": "A different attempt is current", + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelResponse" + } + } + } + } + }, "/api/models/{id}/load-status": { "get": { - "description": "Returns the live state of a distributed cold load — phase, node, byte progress and ETA — or 404 when no load is running for the model. This is the same ` + "`" + `loading` + "`" + ` object the 503 response carries while a model is still staging.", + "description": "Returns the live state of a distributed cold load: job id, phase, node, byte progress, ETA, lease freshness, and for a failed attempt the cause, whether the stop is pending and when the model is released. 404 means no load exists for the model. A database error is 503, never an empty answer. This is the same ` + "`" + `loading` + "`" + ` object the 503 response carries while a model is still staging.", "produces": [ "application/json" ], @@ -1309,6 +1386,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/schema.ErrorResponse" } + }, + "503": { + "description": "The job table could not be read", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } } } } @@ -6884,6 +6967,34 @@ const docTemplate = `{ } } }, + "schema.ModelLoadCancelRequest": { + "type": "object", + "properties": { + "job_id": { + "type": "string" + } + } + }, + "schema.ModelLoadCancelResponse": { + "type": "object", + "properties": { + "current_job_id": { + "type": "string" + }, + "job_id": { + "type": "string" + }, + "model": { + "type": "string" + }, + "retry_after": { + "type": "integer" + }, + "state": { + "type": "string" + } + } + }, "schema.ModelLoadRequest": { "type": "object", "properties": { @@ -6914,6 +7025,10 @@ const docTemplate = `{ "bytes_sent": { "type": "integer" }, + "cancel_requested": { + "description": "CancelRequested is true when an administrator cancelled the attempt.", + "type": "boolean" + }, "eta_seconds": { "description": "ETASeconds is omitted rather than guessed until enough bytes have moved\nfor the observed rate to mean anything. A confidently wrong ETA on a\ntwenty-minute wait is worse than none.", "type": "integer" @@ -6921,6 +7036,18 @@ const docTemplate = `{ "file_index": { "type": "integer" }, + "job_id": { + "description": "JobID names the load attempt. A cancel must quote it: it is the\nprecondition that keeps a cancel from hitting a replacement attempt.", + "type": "string" + }, + "last_error": { + "description": "LastError is the cause of a failed attempt.", + "type": "string" + }, + "lease_expires_in": { + "description": "LeaseExpiresIn is the seconds left on the owner's lease. It is negative\nwhen the lease already ran out, which means the owner is gone.", + "type": "integer" + }, "model": { "type": "string" }, @@ -6930,9 +7057,20 @@ const docTemplate = `{ "progress": { "type": "number" }, + "retry_after": { + "description": "RetryAfter is the seconds until a new load may start, for a failed attempt.", + "type": "integer" + }, "state": { "type": "string" }, + "stop_deadline": { + "type": "string" + }, + "stopping": { + "description": "Stopping is true while the remote work of a failed attempt is not yet\nconfirmed ended. StopDeadline is when the model is released regardless.", + "type": "boolean" + }, "total_bytes": { "type": "integer" }, diff --git a/swagger/swagger.json b/swagger/swagger.json index bff76511f..09676f56c 100644 --- a/swagger/swagger.json +++ b/swagger/swagger.json @@ -1275,9 +1275,86 @@ } } }, + "/api/models/{id}/load-cancel": { + "post": { + "description": "Cancels the load attempt named by `job_id` and stops its remote work. 200 means the attempt is gone or the worker confirmed the stop. 202 means the cancel is recorded and the stop is pending; the model is released after `retry_after` seconds regardless. 409 means a different attempt is current and carries its `current_job_id`. Repeating the call is safe and does not extend the hold.", + "consumes": [ + "application/json" + ], + "produces": [ + "application/json" + ], + "tags": [ + "models" + ], + "summary": "Cancel one distributed model load.", + "parameters": [ + { + "type": "string", + "description": "Model ID", + "name": "id", + "in": "path", + "required": true + }, + { + "description": "The exact attempt to cancel", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelRequest" + } + } + ], + "responses": { + "200": { + "description": "Stopped, or no such load any more", + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelResponse" + } + }, + "202": { + "description": "Cancel recorded; stop pending", + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelResponse" + } + }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "401": { + "description": "Unauthorized", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "403": { + "description": "Forbidden", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "404": { + "description": "Unknown model", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } + }, + "409": { + "description": "A different attempt is current", + "schema": { + "$ref": "#/definitions/schema.ModelLoadCancelResponse" + } + } + } + } + }, "/api/models/{id}/load-status": { "get": { - "description": "Returns the live state of a distributed cold load — phase, node, byte progress and ETA — or 404 when no load is running for the model. This is the same `loading` object the 503 response carries while a model is still staging.", + "description": "Returns the live state of a distributed cold load: job id, phase, node, byte progress, ETA, lease freshness, and for a failed attempt the cause, whether the stop is pending and when the model is released. 404 means no load exists for the model. A database error is 503, never an empty answer. This is the same `loading` object the 503 response carries while a model is still staging.", "produces": [ "application/json" ], @@ -1306,6 +1383,12 @@ "schema": { "$ref": "#/definitions/schema.ErrorResponse" } + }, + "503": { + "description": "The job table could not be read", + "schema": { + "$ref": "#/definitions/schema.ErrorResponse" + } } } } @@ -6881,6 +6964,34 @@ } } }, + "schema.ModelLoadCancelRequest": { + "type": "object", + "properties": { + "job_id": { + "type": "string" + } + } + }, + "schema.ModelLoadCancelResponse": { + "type": "object", + "properties": { + "current_job_id": { + "type": "string" + }, + "job_id": { + "type": "string" + }, + "model": { + "type": "string" + }, + "retry_after": { + "type": "integer" + }, + "state": { + "type": "string" + } + } + }, "schema.ModelLoadRequest": { "type": "object", "properties": { @@ -6911,6 +7022,10 @@ "bytes_sent": { "type": "integer" }, + "cancel_requested": { + "description": "CancelRequested is true when an administrator cancelled the attempt.", + "type": "boolean" + }, "eta_seconds": { "description": "ETASeconds is omitted rather than guessed until enough bytes have moved\nfor the observed rate to mean anything. A confidently wrong ETA on a\ntwenty-minute wait is worse than none.", "type": "integer" @@ -6918,6 +7033,18 @@ "file_index": { "type": "integer" }, + "job_id": { + "description": "JobID names the load attempt. A cancel must quote it: it is the\nprecondition that keeps a cancel from hitting a replacement attempt.", + "type": "string" + }, + "last_error": { + "description": "LastError is the cause of a failed attempt.", + "type": "string" + }, + "lease_expires_in": { + "description": "LeaseExpiresIn is the seconds left on the owner's lease. It is negative\nwhen the lease already ran out, which means the owner is gone.", + "type": "integer" + }, "model": { "type": "string" }, @@ -6927,9 +7054,20 @@ "progress": { "type": "number" }, + "retry_after": { + "description": "RetryAfter is the seconds until a new load may start, for a failed attempt.", + "type": "integer" + }, "state": { "type": "string" }, + "stop_deadline": { + "type": "string" + }, + "stopping": { + "description": "Stopping is true while the remote work of a failed attempt is not yet\nconfirmed ended. StopDeadline is when the model is released regardless.", + "type": "boolean" + }, "total_bytes": { "type": "integer" }, diff --git a/swagger/swagger.yaml b/swagger/swagger.yaml index 244507827..0e2384b9a 100644 --- a/swagger/swagger.yaml +++ b/swagger/swagger.yaml @@ -1793,6 +1793,24 @@ definitions: object: type: string type: object + schema.ModelLoadCancelRequest: + properties: + job_id: + type: string + type: object + schema.ModelLoadCancelResponse: + properties: + current_job_id: + type: string + job_id: + type: string + model: + type: string + retry_after: + type: integer + state: + type: string + type: object schema.ModelLoadRequest: properties: model: @@ -1816,6 +1834,9 @@ definitions: properties: bytes_sent: type: integer + cancel_requested: + description: CancelRequested is true when an administrator cancelled the attempt. + type: boolean eta_seconds: description: |- ETASeconds is omitted rather than guessed until enough bytes have moved @@ -1824,14 +1845,38 @@ definitions: type: integer file_index: type: integer + job_id: + description: |- + JobID names the load attempt. A cancel must quote it: it is the + precondition that keeps a cancel from hitting a replacement attempt. + type: string + last_error: + description: LastError is the cause of a failed attempt. + type: string + lease_expires_in: + description: |- + LeaseExpiresIn is the seconds left on the owner's lease. It is negative + when the lease already ran out, which means the owner is gone. + type: integer model: type: string node: type: string progress: type: number + retry_after: + description: RetryAfter is the seconds until a new load may start, for a failed + attempt. + type: integer state: type: string + stop_deadline: + type: string + stopping: + description: |- + Stopping is true while the remote work of a failed attempt is not yet + confirmed ended. StopDeadline is when the model is released regardless. + type: boolean total_bytes: type: integer total_files: @@ -4194,12 +4239,70 @@ paths: summary: Get an instruction's API guide or OpenAPI fragment tags: - instructions + /api/models/{id}/load-cancel: + post: + consumes: + - application/json + description: Cancels the load attempt named by `job_id` and stops its remote + work. 200 means the attempt is gone or the worker confirmed the stop. 202 + means the cancel is recorded and the stop is pending; the model is released + after `retry_after` seconds regardless. 409 means a different attempt is current + and carries its `current_job_id`. Repeating the call is safe and does not + extend the hold. + parameters: + - description: Model ID + in: path + name: id + required: true + type: string + - description: The exact attempt to cancel + in: body + name: request + required: true + schema: + $ref: '#/definitions/schema.ModelLoadCancelRequest' + produces: + - application/json + responses: + "200": + description: Stopped, or no such load any more + schema: + $ref: '#/definitions/schema.ModelLoadCancelResponse' + "202": + description: Cancel recorded; stop pending + schema: + $ref: '#/definitions/schema.ModelLoadCancelResponse' + "400": + description: Bad Request + schema: + $ref: '#/definitions/schema.ErrorResponse' + "401": + description: Unauthorized + schema: + $ref: '#/definitions/schema.ErrorResponse' + "403": + description: Forbidden + schema: + $ref: '#/definitions/schema.ErrorResponse' + "404": + description: Unknown model + schema: + $ref: '#/definitions/schema.ErrorResponse' + "409": + description: A different attempt is current + schema: + $ref: '#/definitions/schema.ModelLoadCancelResponse' + summary: Cancel one distributed model load. + tags: + - models /api/models/{id}/load-status: get: - description: Returns the live state of a distributed cold load — phase, node, - byte progress and ETA — or 404 when no load is running for the model. This + description: 'Returns the live state of a distributed cold load: job id, phase, + node, byte progress, ETA, lease freshness, and for a failed attempt the cause, + whether the stop is pending and when the model is released. 404 means no load + exists for the model. A database error is 503, never an empty answer. This is the same `loading` object the 503 response carries while a model is still - staging. + staging.' parameters: - description: Model ID in: path @@ -4217,6 +4320,10 @@ paths: description: No load is running for this model schema: $ref: '#/definitions/schema.ErrorResponse' + "503": + description: The job table could not be read + schema: + $ref: '#/definitions/schema.ErrorResponse' summary: Report the progress of an in-flight model load. tags: - models