diff --git a/.agents/distributed-seams.md b/.agents/distributed-seams.md new file mode 100644 index 000000000..a14649d8d --- /dev/null +++ b/.agents/distributed-seams.md @@ -0,0 +1,169 @@ +# Distributed seams + +Distributed mode talks to other processes through a few interfaces. Code above +them must not name NATS, and a new transport is a new implementation of them. +Today every seam has one carrier: NATS, or a direct dial for the gRPC and HTTP +connections to a worker. Only the carrier files import `nats.go` +(`core/services/messaging/client.go`, `tls.go` and +`core/services/nodes/control_nats.go`). + +| Seam | Interface | Lives in | Today | +|---|---|---|---| +| Fan-out | `messaging.Broadcaster` | `core/services/messaging` | NATS client | +| Queues | `messaging.WorkQueue` (producer), `messaging.WorkConsumer` (worker), keyed by `messaging.WorkKind` | `core/services/messaging` (`workqueue.go`) | `NewNATSWorkQueue`, `NewNATSWorkConsumer` | +| Control verbs, frontend half | `nodes.NodeCommandSender`, `nodes.FileStager` | `core/services/nodes` | `RemoteUnloaderAdapter`, `S3NATSFileStager` over NATS request/reply; `HTTPFileStager` over HTTP | +| Control verbs, worker half | unexported `controlServer` (`handle`, `handleWithProgress`) | `core/services/worker` | `natsControlServer` | +| Agent RPC, frontend half | `AgentControl` (package `mcp`) | `core/http/endpoints/mcp` | `nodes.NATSAgentControl` | +| Agent RPC, worker half | unexported `agentRPCServer` | `core/cli/agent_worker.go` | `nodes.NATSAgentRPCServer` | +| Dial to a worker | `nodes.BackendClientFactory`, `nodes.ModelProber`, `nodes.WorkerNetDialerFor` | `core/services/nodes` | direct dial | + +Request and reply payloads of the control verbs live in +`core/services/workerctl`, so the frontend half, the worker half and any +carrier share one wire format without importing each other. + +## Queues + +`WorkKind` names the work (`WorkTask`, `WorkMCPCI`, `WorkAgentRun`). +`natsRoute` in `workqueue_nats.go` is the one place a kind becomes a subject and +a queue group. A nil error from `Enqueue` means the carrier accepted the +payload, not that a consumer exists. The NATS carrier ignores the ctx of +`Enqueue`. + +`Consume(ctx, kind, maxInFlight, h)` keeps two concurrency models on purpose. +With `maxInFlight` 1 the handler runs inline on the NATS delivery goroutine and +a panic is not recovered. Any other value spawns a recovered goroutine per +delivery; when bounded, the slot is taken on the delivery goroutine. 0 and +negative values are unbounded. `Unsubscribe` stops delivery, then waits for +running handlers, so a handler must not call it. The agent worker asks for +MCP CI with 1 (`startMCPCIConsumer`) and for agent runs with the dispatcher's +`maxConcurrent` (0 from the CLI); specs pin both. +`WithAgentRunRoute` lets an agent worker move the agent-run subject and group. +An empty subject keeps `agent.execute`; an empty queue is kept and makes a +plain subscription, as an explicitly empty `LOCALAI_AGENT_QUEUE` always did. + +## Control verbs + +Frontend half: `NodeCommandSender` sends the lifecycle verbs; `FileStager` +moves files. `S3NATSFileStager` returns `nodes.ErrNoRoute` when nothing is +listening for the node. `HTTPFileStager` reports connection failures as +ordinary errors. + +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 +goroutine per request, progress published on the install-progress subject). +An undecodable body is still answered with the verb's typed refusal. The +`undecodable` error a `controlHandler` returns is read only by tests today: the +NATS server drops it. It is a recorded exception to the no-dead-code rule, kept +as the hook a carrier that signals a malformed request out of band (HTTP 400) +needs. + +## Agent RPC + +`AgentControl` carries MCP tool and discovery requests to one agent worker. A +decoded reply is returned with a nil error even when its `Error` field is set. +The NATS implementation reads only the deadline of ctx, never its +cancellation, and uses the default MCP timeouts when there is none. +`NATSAgentRPCServer` serves both in the agent-workers queue group and answers an +undecodable body with an `unmarshal error: ` reply. It also serves the node's +backend stop as `func(backend string)` and never replies to it. + +## Dial + +`BackendClientFactory.NewClient(nodeID, address, parallel)` builds the gRPC +client for SmartRouter, HealthMonitor and the reconciler's default +`ModelProber`. The node id is there because a dialer that must know which node +it reaches, as a tunnel does, cannot recover it from the address; the direct +factory ignores it. `ModelProber.Probe(ctx, nodeID, address)` likewise. + +`WorkerNetDialerFor` returns the dial function for one worker's own HTTP +server. `DirectWorkerNetDialer` dials the address it is handed. It serves +`HTTPFileStager` and the backend-logs proxy (HTTP and WebSocket). +`HTTPFileStager.clientFor` keeps one HTTP client per node, because the idle +pool is keyed by host and port and two workers can report the same address. + +## Rules a carrier must keep + +- Subjects: one closed set of roots, the `broadcastRoots` and `controlRoots` + maps in `core/services/messaging/subject_rules.go`. Add a root there, with a + row in the roots table in `subject_rules_test.go`, not in a carrier. + `messaging.ValidateSubject` refuses anything else with + `messaging.ErrUnservedSubject`. +- Wildcards: only a whole single token `*`, never the root. `>` is refused with + `messaging.ErrUnsupportedWildcard`. +- Delivery is at-most-once. Anything that must survive a gap belongs in a table. +- A subscriber that is reconnecting misses messages. Do not read silence as + evidence about a node. + +## Four conditions that are never reported as each other + +1. A routing fact: no route from here right now (`nodes.ErrNoRoute`). +2. A connection absent within the reconnect grace: nothing acts on it. +3. An unreachable peer: not a verdict about the worker. +4. The worker's own answer, including a refusal: the worker is present. + +On `ErrNoRoute` the only change to the node's own state is the status-only +`MarkUnhealthy`, which the next heartbeat reverses. A caller may also route +around the node (the scheduler skips it, and an upgrade falls back to the older +install subject). Never delete `node_models` rows on it. Timeouts are not +`ErrNoRoute`. + +A pending backend op that fails with `ErrNoRoute` is still recorded as a failed +attempt (`RecordPendingBackendOpFailure`, the reconciler's attempt count, and +the dead-letter after the maximum attempts). That is accounting for the op, not +a verdict about the node. + +The carrier's own sentinel (`nats.ErrNoResponders`) is mapped onto +`ErrNoRoute` in `core/services/nodes/control_nats.go` and nowhere else. A +consumer that matches on the carrier error reads absence as a fact about the +node, which is the mistake this contract exists to prevent. + +## Testing a new carrier + +Run it against `messagingtest.RunBroadcasterConformance` +(`core/services/messaging/messagingtest`) in its own package. `FakeBus` runs it +in `core/services/testutil`. The NATS run is in `core/services/messaging`; it +needs Docker. Without Docker it skips, except when `CI` is set on a non-macOS +runner, where it fails so a runner without Docker cannot hide the check. The +end-to-end specs in `tests/e2e/distributed` run the NATS implementations of the +other seams against a real server, also through Docker. + +## Open items for a second carrier + +- `DistributedModelStore.Range` (`core/services/nodes/distributed_store.go`) + builds a tokenless `model.Model` from `node.Address`. It dials nothing today, + but it is the one direct construction site left outside the dial seam. +- The worker's `files.ensure` handler passes the first caller's ctx into a + shared singleflight closure. The NATS server hands it `context.Background`. + A carrier with a request ctx needs `context.WithoutCancel` there, or one + cancelled caller fails the others. +- `HTTPFileStager` caches a client per node and never forgets one. There is no + `ForgetNode`. +- `clientFor` returns no error; a carrier with no dialer for a node needs that + path. +- The agent worker still uses NATS directly for its connection and for agent + events (`agents.NewEventBridge`). +- `ReplicaReconcilerOptions` has no `ClientFactory` field. Without a + `Prober`, `NewReplicaReconciler` builds its own `tokenClientFactory`, so a + carrier that dials through another factory must add the field. +- The backend-logs proxy (`proxyHTTPToWorker`) starts from + `httpclient.HardenedTransport()`, which keeps + `Proxy: http.ProxyFromEnvironment`. With `HTTP_PROXY` set, a dialer that + routes by node would be handed the proxy address. Clear `Proxy` when the + dialer is not the direct one. +- `natsControlServer.subject` refuses a verb it has no subject for, so + registration fails. A verb that only one carrier serves needs a per-carrier + opt-out where the verbs are registered. +- `WorkHandler` returns only an error. A carrier whose stream handler must send + a terminal reply derives it from the `jobs..result` event the handler + publishes on `events` (`handleMCPCIJob` does this). +- Not additive: the agent-run consumer (`NATSDispatcher.runDelivery`) ignores + the per-delivery `events` publisher and publishes through the process-wide + `EventBridge`, which is built on the NATS client. A second carrier must + change `NATSDispatcher` and `EventBridge`, for example by binding + `handleJob`'s publishes to `events` through a bridge view that shares the + cancel registry. +- `controlHandler`'s `undecodable` return is the recorded exception described + under Control verbs. +- Agent cancel has no production sender (`EventBridge.CancelExecution` has no + caller), so it is not part of the agent RPC seam. diff --git a/AGENTS.md b/AGENTS.md index e3e8b24f7..c34df87dd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -34,6 +34,7 @@ LocalAI follows the Linux kernel project's [guidelines for AI coding assistants] | [.agents/backend-signing.md](.agents/backend-signing.md) | Backend OCI image signing (keyless cosign + sigstore-go) — producer-side CI setup, consumer-side gallery `verification:` block, strict mode (`LOCALAI_REQUIRE_BACKEND_INTEGRITY`), revocation via `not_before` | | [.agents/preparing-a-release.md](.agents/preparing-a-release.md) | Cutting a release: PR labels, `RELEASE_NOTES_vX.Y.Z.md`, the blog post under `website/content/blog/`, and the demo clips under `website/static/media/` | | [.agents/distributed-state.md](.agents/distributed-state.md) | Features that keep runtime state — how they must behave with several frontends (syncstate, advisory-lock leaders, fakebus tests) | +| [.agents/distributed-seams.md](.agents/distributed-seams.md) | Distributed mode transports: the fan-out, queue, control verb, agent RPC and dial seams, subject rules, the no-route contract, conformance suites, open items for a second carrier | | [.impeccable.md](.impeccable.md) | Design context for UI/UX work — users, brand personality, aesthetic direction, and design principles | ## Quick Reference diff --git a/core/application/agent_pool_options_test.go b/core/application/agent_pool_options_test.go new file mode 100644 index 000000000..6cec91e7a --- /dev/null +++ b/core/application/agent_pool_options_test.go @@ -0,0 +1,39 @@ +package application + +import ( + "context" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/messaging" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type stubWorkQueue struct{} + +func (stubWorkQueue) Enqueue(context.Context, messaging.WorkKind, any) error { return nil } + +var _ = Describe("agentPoolOptions", func() { + It("leaves the work queue unset without distributed services", func() { + app := &Application{applicationConfig: &config.ApplicationConfig{}} + + opts := app.agentPoolOptions() + + // Strict comparison: the agent pool reads a non-nil interface as + // distributed mode, and Gomega's BeNil would accept a typed nil. + Expect(opts.WorkQueue == nil).To(BeTrue()) + }) + + It("hands the agent pool the distributed work queue", func() { + queue := stubWorkQueue{} + app := &Application{ + applicationConfig: &config.ApplicationConfig{}, + distributed: &DistributedServices{WorkQueue: queue}, + } + + opts := app.agentPoolOptions() + + Expect(opts.WorkQueue).To(Equal(messaging.WorkQueue(queue))) + }) +}) diff --git a/core/application/application.go b/core/application/application.go index 56e22eeed..7ac7479c5 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -651,14 +651,10 @@ func (a *Application) start() error { return nil } -// StartAgentPool initializes and starts the agent pool service (LocalAGI integration). -// This must be called after the HTTP server is listening, because backends like -// PostgreSQL need to call the embeddings API during collection initialization. -func (a *Application) StartAgentPool() { - if !a.applicationConfig.AgentPool.Enabled { - return - } - // Build options struct from available dependencies +// agentPoolOptions builds the agent pool's dependencies. WorkQueue stays a nil +// interface without distributed services, because the pool reads a non-nil +// WorkQueue as distributed mode. +func (a *Application) agentPoolOptions() agentpool.AgentPoolOptions { opts := agentpool.AgentPoolOptions{ AuthDB: a.authDB, } @@ -666,12 +662,21 @@ func (a *Application) StartAgentPool() { if d.DistStores != nil && d.DistStores.Skills != nil { opts.SkillStore = d.DistStores.Skills } - opts.NATSClient = d.Nats + opts.WorkQueue = d.WorkQueue opts.EventBridge = d.AgentBridge opts.AgentStore = d.AgentStore } + return opts +} - aps, err := agentpool.NewAgentPoolService(a.applicationConfig, opts) +// StartAgentPool initializes and starts the agent pool service (LocalAGI integration). +// This must be called after the HTTP server is listening, because backends like +// PostgreSQL need to call the embeddings API during collection initialization. +func (a *Application) StartAgentPool() { + if !a.applicationConfig.AgentPool.Enabled { + return + } + aps, err := agentpool.NewAgentPoolService(a.applicationConfig, a.agentPoolOptions()) if err != nil { xlog.Error("Failed to create agent pool service", "error", err) return diff --git a/core/application/distributed.go b/core/application/distributed.go index 867743a33..91cb2d053 100644 --- a/core/application/distributed.go +++ b/core/application/distributed.go @@ -11,6 +11,7 @@ import ( "github.com/google/uuid" "github.com/mudler/LocalAI/core/config" + mcpTools "github.com/mudler/LocalAI/core/http/endpoints/mcp" "github.com/mudler/LocalAI/core/services/agents" "github.com/mudler/LocalAI/core/services/distributed" "github.com/mudler/LocalAI/core/services/jobs" @@ -28,6 +29,8 @@ import ( // DistributedServices holds all services initialized for distributed mode. type DistributedServices struct { Nats *messaging.Client + WorkQueue messaging.WorkQueue + AgentControl mcpTools.AgentControl Store storage.ObjectStore Registry *nodes.NodeRegistry Router *nodes.SmartRouter @@ -44,6 +47,10 @@ type DistributedServices struct { Unloader *nodes.RemoteUnloaderAdapter ModelCleanup *nodes.ModelCleanupService + // WorkerHTTPDial reaches a worker's own HTTP server for the admin + // backend-logs proxy, the same way the HTTP file stager does. + WorkerHTTPDial nodes.WorkerNetDialerFor + shutdownOnce sync.Once } @@ -225,8 +232,10 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade } xlog.Info("Distributed job store initialized") + workQueue := messaging.NewNATSWorkQueue(natsClient) + // Initialize job dispatcher - dispatcher := jobs.NewDispatcher(jobStore, natsClient, authDB, cfg.Distributed.InstanceID, cfg.Distributed.JobWorkerConcurrency) + dispatcher := jobs.NewDispatcher(jobStore, workQueue, natsClient, authDB, cfg.Distributed.InstanceID) // Initialize agent store agentStore, err := agents.NewAgentStore(authDB) @@ -261,6 +270,7 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade xlog.Info("File manager initialized", "cacheDir", cacheDir) // Create FileStager for distributed file transfer + workerHTTPDial := nodes.DirectWorkerNetDialer() var fileStager nodes.FileStager if cfg.Distributed.StorageURL != "" { fileStager = nodes.NewS3NATSFileStager(fileMgr, natsClient) @@ -275,7 +285,7 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade return "", fmt.Errorf("node %s has no HTTP address for file transfer", nodeID) } return node.HTTPAddress, nil - }, cfg.Distributed.RegistrationToken) + }, cfg.Distributed.RegistrationToken, workerHTTPDial) xlog.Info("File stager initialized (HTTP direct transfer)") } // Create RemoteUnloaderAdapter — needed by SmartRouter and startup.go @@ -474,6 +484,8 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade success = true return &DistributedServices{ Nats: natsClient, + WorkQueue: workQueue, + AgentControl: nodes.NewNATSAgentControl(natsClient), Store: store, Registry: registry, Router: router, @@ -489,6 +501,8 @@ func initDistributed(cfg *config.ApplicationConfig, authDB *gorm.DB, configLoade ModelAdapter: modelAdapter, Unloader: remoteUnloader, ModelCleanup: modelCleanup, + + WorkerHTTPDial: workerHTTPDial, }, nil } diff --git a/core/application/startup.go b/core/application/startup.go index fe8ee1394..8c8584607 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -324,8 +324,7 @@ func New(opts ...config.AppOption) (*Application, error) { go distSvc.ModelCleanup.Run(options.Context) // In distributed mode, MCP CI jobs are executed by agent workers (not the frontend) // because the frontend can't create MCP sessions (e.g., stdio servers using docker). - // The dispatcher still subscribes to jobs.new for persistence (result/progress subs) - // but does NOT set a workerFn — agent workers consume jobs from the same NATS queue. + // The dispatcher only enqueues jobs and persists the results and traces workers publish. // Wire model config loader so job events include model config for agent workers distSvc.Dispatcher.SetModelConfigLoader(application.backendLoader) diff --git a/core/cli/agent_worker.go b/core/cli/agent_worker.go index 11f515c51..9c5422a51 100644 --- a/core/cli/agent_worker.go +++ b/core/cli/agent_worker.go @@ -19,6 +19,7 @@ import ( "github.com/mudler/LocalAI/core/services/jobs" mcpRemote "github.com/mudler/LocalAI/core/services/mcp" "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/nodes" "github.com/mudler/LocalAI/internal" "github.com/mudler/LocalAI/pkg/sanitize" "github.com/mudler/cogito" @@ -50,7 +51,7 @@ type AgentWorkerCMD struct { APIToken string `env:"LOCALAI_API_TOKEN" help:"API token for LocalAI inference (auto-provisioned during registration if not set)" group:"api"` // NATS subjects - Subject string `env:"LOCALAI_AGENT_SUBJECT" default:"agent.execute" help:"NATS subject for agent execution" group:"distributed"` + Subject string `env:"LOCALAI_AGENT_SUBJECT" default:"agent.execute" help:"NATS subject for agent execution. Must be a served subject (use the agent root, for example agent.execute); an unserved root is refused at startup" group:"distributed"` Queue string `env:"LOCALAI_AGENT_QUEUE" default:"agent-workers" help:"NATS queue group name" group:"distributed"` NatsJWT string `env:"LOCALAI_NATS_JWT" help:"NATS user JWT override (defaults to nats_jwt from registration)" group:"distributed"` @@ -75,7 +76,22 @@ func (cmd *AgentWorkerCMD) natsAuthRequired() bool { return cmd.NatsRequireAuth || cmd.DistributedRequireAuth } +// validateAgentSubject refuses a subject no carrier serves before the worker +// registers or connects. The messaging client refuses it anyway at subscribe +// time, but by then the worker has registered and the error does not name the +// setting the operator has to change. +func validateAgentSubject(subject string) error { + if err := messaging.ValidateSubject(subject); err != nil { + return fmt.Errorf("LOCALAI_AGENT_SUBJECT %q must be a served subject (use the agent root, for example %s): %w", + subject, messaging.SubjectAgentExecute, err) + } + return nil +} + func (cmd *AgentWorkerCMD) Run(ctx *cliContext.Context) error { + if err := validateAgentSubject(cmd.Subject); err != nil { + return err + } xlog.Info("Starting agent worker", "nats", sanitize.URL(cmd.NatsURL), "register_to", cmd.RegisterTo) // Resolve API URL @@ -184,14 +200,17 @@ func (cmd *AgentWorkerCMD) Run(ctx *cliContext.Context) error { defer cancelSub.Unsubscribe() } + // One consumer serves both queued kinds; the route option only moves the + // agent-run subject and group, which operators may set. + work := messaging.NewNATSWorkConsumer(natsClient, messaging.WithAgentRunRoute(cmd.Subject, cmd.Queue)) + // Create and start the NATS dispatcher. // No ConfigProvider or SkillStore needed — config and skills arrive in the job payload. dispatcher := agents.NewNATSDispatcher( - natsClient, + work, eventBridge, nil, // no ConfigProvider: config comes in the enriched NATS payload apiURL, cmd.APIToken, - cmd.Subject, cmd.Queue, 0, // no concurrency limit (CLI worker) ) @@ -199,19 +218,17 @@ func (cmd *AgentWorkerCMD) Run(ctx *cliContext.Context) error { return fmt.Errorf("starting dispatcher: %w", err) } - // Subscribe to MCP tool execution requests (load-balanced across workers). - // The frontend routes model-level MCP tool calls here via NATS request-reply. - if _, err := natsClient.QueueSubscribeReply(messaging.SubjectMCPToolExecute, messaging.QueueAgentWorkers, func(data []byte, reply func([]byte)) { - handleMCPToolRequest(data, reply) - }); err != nil { - return fmt.Errorf("subscribing to %s: %w", messaging.SubjectMCPToolExecute, err) + var rpc agentRPCServer = nodes.NewNATSAgentRPCServer(natsClient, nodeID) + + // Serve MCP tool execution requests (load-balanced across workers). + // The frontend routes model-level MCP tool calls here. + if err := rpc.ServeMCPTool(handleMCPToolRequest); err != nil { + return err } - // Subscribe to MCP discovery requests (load-balanced across workers). - if _, err := natsClient.QueueSubscribeReply(messaging.SubjectMCPDiscovery, messaging.QueueAgentWorkers, func(data []byte, reply func([]byte)) { - handleMCPDiscoveryRequest(data, reply) - }); err != nil { - return fmt.Errorf("subscribing to %s: %w", messaging.SubjectMCPDiscovery, err) + // Serve MCP discovery requests (load-balanced across workers). + if err := rpc.ServeMCPDiscovery(handleMCPDiscoveryRequest); err != nil { + return err } // Subscribe to MCP CI job execution (load-balanced across agent workers). @@ -223,24 +240,19 @@ func (cmd *AgentWorkerCMD) Run(ctx *cliContext.Context) error { } mcpCIJobTimeout = cmp.Or(mcpCIJobTimeout, config.DefaultMCPCIJobTimeout) - if _, err := natsClient.QueueSubscribe(messaging.SubjectMCPCIJobsNew, messaging.QueueWorkers, func(data []byte) { - handleMCPCIJob(shutdownCtx, data, apiURL, cmd.APIToken, natsClient, mcpCIJobTimeout) - }); err != nil { - return fmt.Errorf("subscribing to %s: %w", messaging.SubjectMCPCIJobsNew, err) + if _, err := startMCPCIConsumer(shutdownCtx, work, apiURL, cmd.APIToken, mcpCIJobTimeout); err != nil { + return err } - // Subscribe to backend stop events to clean up cached MCP sessions. + // Listen for backend stop events to clean up cached MCP sessions. // In the main application this is done via ml.OnModelUnload, but the agent - // worker has no model loader — we listen for the NATS stop event instead. - if _, err := natsClient.Subscribe(messaging.SubjectNodeBackendStop(nodeID), func(data []byte) { - var req struct { - Backend string `json:"backend"` - } - if json.Unmarshal(data, &req) == nil && req.Backend != "" { - mcpTools.CloseMCPSessions(req.Backend) + // worker has no model loader, so it listens for the stop event instead. + if err := rpc.ServeBackendStop(func(backend string) { + if backend != "" { + mcpTools.CloseMCPSessions(backend) } }); err != nil { - return fmt.Errorf("subscribing to %s: %w", messaging.SubjectNodeBackendStop(nodeID), err) + return err } xlog.Info("Agent worker ready, waiting for jobs", "subject", cmd.Subject, "queue", cmd.Queue) @@ -266,72 +278,65 @@ func (cmd *AgentWorkerCMD) Run(ctx *cliContext.Context) error { return runErr } -// handleMCPToolRequest handles a NATS request-reply for MCP tool execution. -// The worker creates/caches MCP sessions from the serialized config and executes the tool. -func handleMCPToolRequest(data []byte, reply func([]byte)) { - var req mcpRemote.MCPToolRequest - if err := json.Unmarshal(data, &req); err != nil { - sendMCPToolReply(reply, "", fmt.Sprintf("unmarshal error: %v", err)) - return - } +// startMCPCIConsumer serves MCP CI jobs with maxInFlight 1, which keeps them +// one at a time per worker, run inline on the delivery, as they always were. +func startMCPCIConsumer(ctx context.Context, consumer messaging.WorkConsumer, apiURL, apiToken string, jobTimeout time.Duration) (messaging.Subscription, error) { + return consumer.Consume(ctx, messaging.WorkMCPCI, 1, func(ctx context.Context, data []byte, events messaging.Publisher) error { + return handleMCPCIJob(ctx, data, apiURL, apiToken, events, jobTimeout) + }) +} - ctx, cancel := context.WithTimeout(context.Background(), config.DefaultMCPToolTimeout) +// agentRPCServer is how the agent worker serves the frontend's MCP requests +// and hears the node's backend stop events. +type agentRPCServer interface { + ServeMCPTool(h mcpRemote.ToolHandler) error + ServeMCPDiscovery(h mcpRemote.DiscoveryHandler) error + ServeBackendStop(h func(backend string)) error +} + +// handleMCPToolRequest executes one MCP tool call. The worker creates/caches +// MCP sessions from the serialized config and executes the tool. +func handleMCPToolRequest(parent context.Context, req mcpRemote.MCPToolRequest) mcpRemote.MCPToolResponse { + ctx, cancel := context.WithTimeout(parent, config.DefaultMCPToolTimeout) defer cancel() // Create/cache named MCP sessions from the provided config namedSessions, err := mcpTools.NamedSessionsFromMCPConfig(req.ModelName, req.RemoteServers, req.StdioServers, nil) if err != nil { - sendMCPToolReply(reply, "", fmt.Sprintf("session error: %v", err)) - return + return mcpRemote.MCPToolResponse{Error: fmt.Sprintf("session error: %v", err)} } // Discover tools to find the right session tools, err := mcpTools.DiscoverMCPTools(ctx, namedSessions) if err != nil { - sendMCPToolReply(reply, "", fmt.Sprintf("discovery error: %v", err)) - return + return mcpRemote.MCPToolResponse{Error: fmt.Sprintf("discovery error: %v", err)} } // Execute the tool argsJSON, _ := json.Marshal(req.Arguments) result, err := mcpTools.ExecuteMCPToolCall(ctx, tools, req.ToolName, string(argsJSON)) if err != nil { - sendMCPToolReply(reply, "", err.Error()) - return + return mcpRemote.MCPToolResponse{Error: err.Error()} } - sendMCPToolReply(reply, result, "") + return mcpRemote.MCPToolResponse{Result: result} } -func sendMCPToolReply(reply func([]byte), result, errMsg string) { - resp := mcpRemote.MCPToolResponse{Result: result, Error: errMsg} - data, _ := json.Marshal(resp) - reply(data) -} - -// handleMCPDiscoveryRequest handles a NATS request-reply for MCP tool/prompt/resource discovery. -func handleMCPDiscoveryRequest(data []byte, reply func([]byte)) { - var req mcpRemote.MCPDiscoveryRequest - if err := json.Unmarshal(data, &req); err != nil { - sendMCPDiscoveryReply(reply, nil, nil, fmt.Sprintf("unmarshal error: %v", err)) - return - } - - ctx, cancel := context.WithTimeout(context.Background(), config.DefaultMCPDiscoveryTimeout) +// handleMCPDiscoveryRequest lists a model's MCP tools, prompts and resources. +func handleMCPDiscoveryRequest(parent context.Context, req mcpRemote.MCPDiscoveryRequest) mcpRemote.MCPDiscoveryResponse { + ctx, cancel := context.WithTimeout(parent, config.DefaultMCPDiscoveryTimeout) defer cancel() // Create/cache named MCP sessions namedSessions, err := mcpTools.NamedSessionsFromMCPConfig(req.ModelName, req.RemoteServers, req.StdioServers, nil) if err != nil { - sendMCPDiscoveryReply(reply, nil, nil, fmt.Sprintf("session error: %v", err)) - return + return mcpRemote.MCPDiscoveryResponse{Error: fmt.Sprintf("session error: %v", err)} } // List servers with their tools/prompts/resources serverInfos, err := mcpTools.ListMCPServers(ctx, namedSessions) if err != nil { - sendMCPDiscoveryReply(reply, nil, nil, fmt.Sprintf("list error: %v", err)) - return + return mcpRemote.MCPDiscoveryResponse{Error: fmt.Sprintf("list error: %v", err)} } // Also get tool function schemas for the frontend @@ -358,55 +363,51 @@ func handleMCPDiscoveryRequest(data []byte, reply func([]byte)) { }) } - sendMCPDiscoveryReply(reply, servers, toolDefs, "") -} - -func sendMCPDiscoveryReply(reply func([]byte), servers []mcpRemote.MCPServerInfo, tools []mcpRemote.MCPToolDef, errMsg string) { - resp := mcpRemote.MCPDiscoveryResponse{Servers: servers, Tools: tools, Error: errMsg} - data, _ := json.Marshal(resp) - reply(data) + return mcpRemote.MCPDiscoveryResponse{Servers: servers, Tools: toolDefs} } // handleMCPCIJob processes an MCP CI job on the agent worker. // The agent worker can create MCP sessions (has docker) and call the LocalAI API for inference. -func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken string, natsClient messaging.MessagingClient, jobTimeout time.Duration) { +// Every outcome, failures included, is reported on events or logged, so it +// always returns nil: a carrier that redelivers on error would only repeat it. +func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken string, events messaging.Publisher, jobTimeout time.Duration) error { var evt jobs.JobEvent if err := json.Unmarshal(data, &evt); err != nil { xlog.Error("Failed to unmarshal job event", "error", err) - return + return nil } job := evt.Job task := evt.Task if job == nil || task == nil { xlog.Error("MCP CI job missing enriched data", "jobID", evt.JobID) - publishJobResult(natsClient, evt.JobID, "failed", "", "job or task data missing from NATS event") - return + publishJobResult(events, evt.JobID, "failed", "", "job or task data missing from NATS event") + return nil } modelCfg := evt.ModelConfig if modelCfg == nil { - publishJobResult(natsClient, evt.JobID, "failed", "", "model config missing from job event") - return + publishJobResult(events, evt.JobID, "failed", "", "model config missing from job event") + return nil } xlog.Info("Processing MCP CI job", "jobID", evt.JobID, "taskID", evt.TaskID, "model", task.Model) // Publish running status - natsClient.Publish(messaging.SubjectJobProgress(evt.JobID), jobs.ProgressEvent{ + publishJobTrace(events, jobs.ProgressEvent{ JobID: evt.JobID, Status: "running", Message: "Job started on agent worker", }) // Parse MCP config if modelCfg.MCP.Servers == "" && modelCfg.MCP.Stdio == "" { - publishJobResult(natsClient, evt.JobID, "failed", "", "no MCP servers configured for model") - return + publishJobResult(events, evt.JobID, "failed", "", "no MCP servers configured for model") + return nil } remote, stdio, err := modelCfg.MCP.MCPConfigFromYAML() if err != nil { - publishJobResult(natsClient, evt.JobID, "failed", "", fmt.Sprintf("failed to parse MCP config: %v", err)) - return + publishJobResult(events, evt.JobID, "failed", "", fmt.Sprintf("failed to parse MCP config: %v", err)) + return nil } // Create MCP sessions locally (agent worker has docker) @@ -416,8 +417,8 @@ func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken s if err != nil { errMsg = fmt.Sprintf("failed to create MCP sessions: %v", err) } - publishJobResult(natsClient, evt.JobID, "failed", "", errMsg) - return + publishJobResult(events, evt.JobID, "failed", "", errMsg) + return nil } // Build prompt from template @@ -449,7 +450,7 @@ func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken s defer cancel() // Update job status to running in DB - publishJobStatus(natsClient, evt.JobID, "running", "") + publishJobStatus(events, evt.JobID, "running", "") // Buffer stream tokens and flush as complete blocks var reasoningBuf, contentBuf strings.Builder @@ -457,13 +458,13 @@ func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken s flushStreamBuf := func() { if reasoningBuf.Len() > 0 { - natsClient.Publish(messaging.SubjectJobProgress(evt.JobID), jobs.ProgressEvent{ + publishJobTrace(events, jobs.ProgressEvent{ JobID: evt.JobID, TraceType: "reasoning", TraceContent: reasoningBuf.String(), }) reasoningBuf.Reset() } if contentBuf.Len() > 0 { - natsClient.Publish(messaging.SubjectJobProgress(evt.JobID), jobs.ProgressEvent{ + publishJobTrace(events, jobs.ProgressEvent{ JobID: evt.JobID, TraceType: "content", TraceContent: contentBuf.String(), }) contentBuf.Reset() @@ -476,13 +477,13 @@ func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken s cogito.WithMCPs(sessions...), cogito.WithStatusCallback(func(status string) { flushStreamBuf() - natsClient.Publish(messaging.SubjectJobProgress(evt.JobID), jobs.ProgressEvent{ + publishJobTrace(events, jobs.ProgressEvent{ JobID: evt.JobID, TraceType: "status", TraceContent: status, }) }), cogito.WithToolCallResultCallback(func(t cogito.ToolStatus) { flushStreamBuf() - natsClient.Publish(messaging.SubjectJobProgress(evt.JobID), jobs.ProgressEvent{ + publishJobTrace(events, jobs.ProgressEvent{ JobID: evt.JobID, TraceType: "tool_result", TraceContent: fmt.Sprintf("%s: %s", t.Name, t.Result), }) }), @@ -498,7 +499,7 @@ func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken s case cogito.StreamEventContent: contentBuf.WriteString(ev.Content) case cogito.StreamEventToolCall: - natsClient.Publish(messaging.SubjectJobProgress(evt.JobID), jobs.ProgressEvent{ + publishJobTrace(events, jobs.ProgressEvent{ JobID: evt.JobID, TraceType: "tool_call", TraceContent: fmt.Sprintf("%s(%s)", ev.ToolName, ev.ToolArgs), }) } @@ -513,22 +514,31 @@ func handleMCPCIJob(shutdownCtx context.Context, data []byte, apiURL, apiToken s flushStreamBuf() // flush any remaining buffered tokens if err != nil { - publishJobResult(natsClient, evt.JobID, "failed", "", fmt.Sprintf("cogito execution failed: %v", err)) - return + publishJobResult(events, evt.JobID, "failed", "", fmt.Sprintf("cogito execution failed: %v", err)) + return nil } result := "" if msg := f.LastMessage(); msg != nil { result = msg.Content } - publishJobResult(natsClient, evt.JobID, "completed", result, "") + publishJobResult(events, evt.JobID, "completed", result, "") xlog.Info("MCP CI job completed", "jobID", evt.JobID, "resultLen", len(result)) + return nil } -func publishJobStatus(nc messaging.MessagingClient, jobID, status, message string) { - jobs.PublishJobProgress(nc, jobID, status, message) +func publishJobStatus(events messaging.Publisher, jobID, status, message string) { + jobs.PublishJobProgress(events, jobID, status, message) } -func publishJobResult(nc messaging.MessagingClient, jobID, status, result, errMsg string) { - jobs.PublishJobResult(nc, jobID, status, result, errMsg) +func publishJobResult(events messaging.Publisher, jobID, status, result, errMsg string) { + jobs.PublishJobResult(events, jobID, status, result, errMsg) +} + +// publishJobTrace sends a progress or trace line; a lost line must not fail +// the job, so the error is only logged. +func publishJobTrace(events messaging.Publisher, ev jobs.ProgressEvent) { + if err := events.Publish(messaging.SubjectJobProgress(ev.JobID), ev); err != nil { + xlog.Error("Failed to publish job progress", "jobID", ev.JobID, "error", err) + } } diff --git a/core/cli/agent_worker_mcp_rpc_test.go b/core/cli/agent_worker_mcp_rpc_test.go new file mode 100644 index 000000000..38b117ed8 --- /dev/null +++ b/core/cli/agent_worker_mcp_rpc_test.go @@ -0,0 +1,47 @@ +package cli + +import ( + "context" + + "github.com/mudler/LocalAI/core/config" + mcpTools "github.com/mudler/LocalAI/core/http/endpoints/mcp" + mcpRemote "github.com/mudler/LocalAI/core/services/mcp" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// The frontend hands a worker's reply to the model as the tool result, so a +// failure has to come back as a reply with Error set: silence would leave the +// caller waiting for the full request budget. +var _ = Describe("agent worker MCP handlers", func() { + var tool mcpRemote.ToolHandler = handleMCPToolRequest + var discovery mcpRemote.DiscoveryHandler = handleMCPDiscoveryRequest + + const model = "agent-worker-mcp-rpc-spec" + + AfterEach(func() { + mcpTools.CloseMCPSessions(model) + }) + + It("answers a tool no server provides with an error reply", func() { + resp := tool(context.Background(), mcpRemote.MCPToolRequest{ModelName: model, ToolName: "missing"}) + Expect(resp.Result).To(BeEmpty()) + Expect(resp.Error).To(ContainSubstring(`MCP tool "missing" not found`)) + }) + + It("answers a discovery whose server cannot start with that server's error", func() { + resp := discovery(context.Background(), mcpRemote.MCPDiscoveryRequest{ + ModelName: model, + StdioServers: config.MCPGenericConfig[config.MCPSTDIOServers]{ + Servers: config.MCPSTDIOServers{ + "broken": {Command: "/nonexistent/localai-spec-mcp-server"}, + }, + }, + }) + Expect(resp.Servers).To(HaveLen(1)) + Expect(resp.Servers[0].Name).To(Equal("broken")) + Expect(resp.Servers[0].Error).To(ContainSubstring("startup failed")) + Expect(resp.Tools).To(BeEmpty()) + }) +}) diff --git a/core/cli/agent_worker_mcpci_test.go b/core/cli/agent_worker_mcpci_test.go new file mode 100644 index 000000000..5c0034da4 --- /dev/null +++ b/core/cli/agent_worker_mcpci_test.go @@ -0,0 +1,139 @@ +package cli + +import ( + "context" + "encoding/json" + "sync" + "time" + + "github.com/mudler/LocalAI/core/services/jobs" + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/testutil" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("handleMCPCIJob", func() { + // The handler must report on the publisher the carrier hands it: a + // carrier without a process-wide bus has no other place to send the + // terminal result. + It("publishes a failed result on the events publisher when the model config is missing", func() { + events := testutil.NewFakeBus() + + var mu sync.Mutex + var results []jobs.JobResultEvent + _, err := events.Subscribe(messaging.SubjectJobResult("job-1"), func(data []byte) { + var r jobs.JobResultEvent + Expect(json.Unmarshal(data, &r)).To(Succeed()) + mu.Lock() + results = append(results, r) + mu.Unlock() + }) + Expect(err).ToNot(HaveOccurred()) + + payload, err := json.Marshal(jobs.JobEvent{ + JobID: "job-1", + TaskID: "task-1", + Job: &jobs.JobRecord{ID: "job-1"}, + Task: &jobs.TaskRecord{ID: "task-1", Model: "m"}, + }) + Expect(err).ToNot(HaveOccurred()) + + Expect(handleMCPCIJob(context.Background(), payload, "http://127.0.0.1:1", "", events, time.Second)).To(Succeed()) + + mu.Lock() + defer mu.Unlock() + Expect(results).To(ConsistOf(jobs.JobResultEvent{ + JobID: "job-1", + Status: "failed", + Error: "model config missing from job event", + })) + }) + + It("drops an undecodable payload without publishing anything", func() { + events := testutil.NewFakeBus() + var mu sync.Mutex + published := 0 + // The subject helpers sanitise "*", so the wildcards are spelled out. + for _, subject := range []string{"jobs.*.result", "jobs.*.progress"} { + _, err := events.Subscribe(subject, func([]byte) { + mu.Lock() + published++ + mu.Unlock() + }) + Expect(err).ToNot(HaveOccurred()) + } + + Expect(handleMCPCIJob(context.Background(), []byte("not json"), "http://127.0.0.1:1", "", events, time.Second)).To(Succeed()) + + mu.Lock() + defer mu.Unlock() + Expect(published).To(BeZero()) + }) +}) + +// recordingWorkConsumer keeps what Consume was asked for, so the spec pins the +// limit the agent worker chooses, not what a carrier does with it. +type recordingWorkConsumer struct { + kind messaging.WorkKind + maxInFlight int + handler messaging.WorkHandler + calls int +} + +func (c *recordingWorkConsumer) Consume(_ context.Context, kind messaging.WorkKind, maxInFlight int, h messaging.WorkHandler) (messaging.Subscription, error) { + c.calls++ + c.kind, c.maxInFlight, c.handler = kind, maxInFlight, h + return nil, nil +} + +var _ = Describe("startMCPCIConsumer", func() { + // MCP CI jobs may start stdio servers in containers; running them one at a + // time per agent worker is the behaviour operators size workers for. + It("consumes MCP CI jobs one at a time", func() { + consumer := &recordingWorkConsumer{} + _, err := startMCPCIConsumer(GinkgoT().Context(), consumer, "http://127.0.0.1:1", "", time.Second) + Expect(err).ToNot(HaveOccurred()) + + Expect(consumer.calls).To(Equal(1)) + Expect(consumer.kind).To(Equal(messaging.WorkMCPCI)) + Expect(consumer.maxInFlight).To(Equal(1)) + }) + + It("runs each delivery through handleMCPCIJob on the delivery's events publisher", func() { + consumer := &recordingWorkConsumer{} + _, err := startMCPCIConsumer(GinkgoT().Context(), consumer, "http://127.0.0.1:1", "", time.Second) + Expect(err).ToNot(HaveOccurred()) + Expect(consumer.handler).ToNot(BeNil()) + + events := testutil.NewFakeBus() + var mu sync.Mutex + var results []jobs.JobResultEvent + _, err = events.Subscribe(messaging.SubjectJobResult("job-2"), func(data []byte) { + var r jobs.JobResultEvent + Expect(json.Unmarshal(data, &r)).To(Succeed()) + mu.Lock() + results = append(results, r) + mu.Unlock() + }) + Expect(err).ToNot(HaveOccurred()) + + payload, err := json.Marshal(jobs.JobEvent{ + JobID: "job-2", + TaskID: "task-2", + Job: &jobs.JobRecord{ID: "job-2"}, + Task: &jobs.TaskRecord{ID: "task-2", Model: "m"}, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(consumer.handler(context.Background(), payload, events)).To(Succeed()) + + mu.Lock() + defer mu.Unlock() + Expect(results).To(ConsistOf(jobs.JobResultEvent{ + JobID: "job-2", + Status: "failed", + Error: "model config missing from job event", + })) + }) +}) diff --git a/core/cli/agent_worker_subject_test.go b/core/cli/agent_worker_subject_test.go new file mode 100644 index 000000000..3f5f14fa3 --- /dev/null +++ b/core/cli/agent_worker_subject_test.go @@ -0,0 +1,31 @@ +package cli + +import ( + "errors" + + "github.com/mudler/LocalAI/core/services/messaging" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("validateAgentSubject", func() { + It("accepts the default agent execution subject", func() { + Expect(validateAgentSubject("agent.execute")).To(Succeed()) + }) + + It("refuses a subject whose root no carrier serves and names the env var", func() { + err := validateAgentSubject("tenant-a.agent.execute") + Expect(err).To(HaveOccurred()) + Expect(errors.Is(err, messaging.ErrUnservedSubject)).To(BeTrue()) + Expect(err.Error()).To(ContainSubstring("LOCALAI_AGENT_SUBJECT")) + Expect(err.Error()).To(ContainSubstring("tenant-a.agent.execute")) + Expect(err.Error()).To(ContainSubstring("agent.execute")) + }) + + It("refuses a multi-token wildcard", func() { + err := validateAgentSubject("agent.>") + Expect(errors.Is(err, messaging.ErrUnsupportedWildcard)).To(BeTrue()) + Expect(err.Error()).To(ContainSubstring("LOCALAI_AGENT_SUBJECT")) + }) +}) diff --git a/core/config/distributed_config.go b/core/config/distributed_config.go index ef7e01bfd..a6e11ffde 100644 --- a/core/config/distributed_config.go +++ b/core/config/distributed_config.go @@ -97,7 +97,6 @@ type DistributedConfig struct { MaxUploadSize int64 // Maximum upload body size in bytes (default 50 GB) AgentWorkerConcurrency int `yaml:"agent_worker_concurrency" json:"agent_worker_concurrency" env:"LOCALAI_AGENT_WORKER_CONCURRENCY"` - JobWorkerConcurrency int `yaml:"job_worker_concurrency" json:"job_worker_concurrency" env:"LOCALAI_JOB_WORKER_CONCURRENCY"` // DiskHeadroomDisabled turns off the scheduler's free-disk admission check, // restoring the pre-#11054 behaviour where node selection ignores whether a diff --git a/core/http/app.go b/core/http/app.go index 4543a0048..b6211686b 100644 --- a/core/http/app.go +++ b/core/http/app.go @@ -561,15 +561,17 @@ func API(application *application.Application) (*echo.Echo, error) { distCfg := application.ApplicationConfig().Distributed var registry *nodes.NodeRegistry var remoteUnloader nodes.NodeCommandSender + var workerHTTPDial nodes.WorkerNetDialerFor if d := application.Distributed(); d != nil { registry = d.Registry + workerHTTPDial = d.WorkerHTTPDial if d.Router != nil { remoteUnloader = d.Router.Unloader() } } natsCfg := distCfg.NatsAuthConfig() routes.RegisterNodeSelfServiceRoutes(e, registry, distCfg.RegistrationToken, distCfg.AutoApproveNodes, application.AuthDB(), application.ApplicationConfig().Auth.APIKeyHMACSecret, natsCfg) - routes.RegisterNodeAdminRoutes(e, registry, remoteUnloader, application.GalleryService(), opcache, application.ApplicationConfig(), adminMiddleware, application.AuthDB(), application.ApplicationConfig().Auth.APIKeyHMACSecret, application.ApplicationConfig().Distributed.RegistrationToken, natsCfg) + routes.RegisterNodeAdminRoutes(e, registry, remoteUnloader, application.GalleryService(), opcache, application.ApplicationConfig(), adminMiddleware, application.AuthDB(), application.ApplicationConfig().Auth.APIKeyHMACSecret, application.ApplicationConfig().Distributed.RegistrationToken, natsCfg, workerHTTPDial) // Distributed SSE routes (job progress + agent events via NATS) if d := application.Distributed(); d != nil { diff --git a/core/http/endpoints/anthropic/messages.go b/core/http/endpoints/anthropic/messages.go index 9310cabd9..231444ded 100644 --- a/core/http/endpoints/anthropic/messages.go +++ b/core/http/endpoints/anthropic/messages.go @@ -29,7 +29,7 @@ import ( // @Param request body schema.AnthropicRequest true "query params" // @Success 200 {object} schema.AnthropicResponse "Response" // @Router /v1/messages [post] -func MessagesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient) echo.HandlerFunc { +func MessagesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, agentControl mcpTools.AgentControl) echo.HandlerFunc { return func(c echo.Context) error { id := uuid.New().String() @@ -70,7 +70,7 @@ func MessagesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evalu if (len(mcpServers) > 0 || mcpPromptName != "" || len(mcpResourceURIs) > 0) && (cfg.MCP.Servers != "" || cfg.MCP.Stdio != "") { remote, stdio, mcpErr := cfg.MCP.MCPConfigFromYAML() if mcpErr == nil { - mcpExecutor = mcpTools.NewToolExecutor(c.Request().Context(), natsClient, cfg.Name, remote, stdio, mcpServers) + mcpExecutor = mcpTools.NewToolExecutor(c.Request().Context(), agentControl, cfg.Name, remote, stdio, mcpServers) // Prompt and resource injection (pre-processing step — resolves locally regardless of distributed mode) namedSessions, sessErr := mcpTools.NamedSessionsFromMCPConfig(cfg.Name, remote, stdio, mcpServers) diff --git a/core/http/endpoints/localai/mcp.go b/core/http/endpoints/localai/mcp.go index f3905442d..9687afac6 100644 --- a/core/http/endpoints/localai/mcp.go +++ b/core/http/endpoints/localai/mcp.go @@ -57,7 +57,7 @@ type MCPErrorEvent struct { // @Param request body schema.OpenAIRequest true "query params" // @Success 200 {object} schema.OpenAIResponse "Response" // @Router /v1/mcp/chat/completions [post] -func MCPEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient, compressor middleware.ChatCompressor) echo.HandlerFunc { +func MCPEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, agentControl mcpTools.AgentControl, compressor middleware.ChatCompressor) echo.HandlerFunc { // The legacy /v1/mcp/chat/completions endpoint never opts into the // in-process LocalAI Assistant tool surface — pass nil holder so the // assistant branch in chat.go is unreachable from this code path. @@ -65,7 +65,7 @@ func MCPEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator // the per-model PII config and is kept for backward compatibility. // The request-side middleware on the main chat route handles // filtering for the standard /v1/chat/completions path. - chatHandler := openai.ChatEndpoint(cl, ml, evaluator, appConfig, natsClient, nil, compressor) + chatHandler := openai.ChatEndpoint(cl, ml, evaluator, appConfig, agentControl, nil, compressor) return func(c echo.Context) error { input, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) diff --git a/core/http/endpoints/localai/mcp_tools.go b/core/http/endpoints/localai/mcp_tools.go index f5db27bd7..8839cd74d 100644 --- a/core/http/endpoints/localai/mcp_tools.go +++ b/core/http/endpoints/localai/mcp_tools.go @@ -13,7 +13,7 @@ import ( // MCPServersEndpoint returns the list of MCP servers and their tools for a given model. // GET /v1/mcp/servers/:model -func MCPServersEndpoint(cl *config.ModelConfigLoader, appConfig *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient) echo.HandlerFunc { +func MCPServersEndpoint(cl *config.ModelConfigLoader, appConfig *config.ApplicationConfig, agentControl mcpTools.AgentControl) echo.HandlerFunc { return func(c echo.Context) error { modelName := c.Param("model") if modelName == "" { @@ -47,8 +47,8 @@ func MCPServersEndpoint(cl *config.ModelConfigLoader, appConfig *config.Applicat // In distributed mode, route discovery through NATS to an agent worker // that can actually connect to the MCP servers. - if natsClient != nil { - resp, err := mcpTools.DiscoverMCPToolsRemote(c.Request().Context(), natsClient, cfg.Name, remote, stdio) + if agentControl != nil { + resp, err := mcpTools.DiscoverMCPToolsRemote(c.Request().Context(), agentControl, cfg.Name, remote, stdio) if err != nil { return c.JSON(http.StatusOK, map[string]any{ "model": modelName, @@ -80,7 +80,7 @@ func MCPServersEndpoint(cl *config.ModelConfigLoader, appConfig *config.Applicat // MCPServersEndpointFromMiddleware is a version that uses the middleware-resolved model config. // This allows it to use the same middleware chain as other endpoints. -func MCPServersEndpointFromMiddleware(natsClient mcpTools.MCPNATSClient) echo.HandlerFunc { +func MCPServersEndpointFromMiddleware(agentControl mcpTools.AgentControl) echo.HandlerFunc { return func(c echo.Context) error { cfg, ok := c.Get(middleware.CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) if !ok || cfg == nil { @@ -104,8 +104,8 @@ func MCPServersEndpointFromMiddleware(natsClient mcpTools.MCPNATSClient) echo.Ha } // In distributed mode, route discovery through NATS to an agent worker. - if natsClient != nil { - resp, err := mcpTools.DiscoverMCPToolsRemote(c.Request().Context(), natsClient, cfg.Name, remote, stdio) + if agentControl != nil { + resp, err := mcpTools.DiscoverMCPToolsRemote(c.Request().Context(), agentControl, cfg.Name, remote, stdio) if err != nil { return c.JSON(http.StatusOK, map[string]any{ "model": cfg.Name, diff --git a/core/http/endpoints/localai/nodes.go b/core/http/endpoints/localai/nodes.go index 2036f3d6e..32c0ffe4e 100644 --- a/core/http/endpoints/localai/nodes.go +++ b/core/http/endpoints/localai/nodes.go @@ -25,9 +25,9 @@ import ( "github.com/mudler/LocalAI/core/http/auth" "github.com/mudler/LocalAI/core/schema" "github.com/mudler/LocalAI/core/services/galleryop" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes" "github.com/mudler/LocalAI/core/services/nodes/prefixcache" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/httpclient" "github.com/mudler/LocalAI/pkg/natsauth" "github.com/mudler/LocalAI/pkg/vrambudget" @@ -668,7 +668,7 @@ func ListBackendsOnNodeEndpoint(unloader nodes.NodeCommandSender, registry *node // single-node and cluster-wide views stay consistent. if node, err := registry.Get(c.Request().Context(), nodeID); err == nil { if node.NodeType != "" && node.NodeType != nodes.NodeTypeBackend { - return c.JSON(http.StatusOK, []messaging.NodeBackendInfo{}) + return c.JSON(http.StatusOK, []workerctl.NodeBackendInfo{}) } } if unloader == nil { @@ -743,7 +743,7 @@ func DeleteModelOnNodeEndpoint(unloader nodes.NodeCommandSender, registry *nodes // NodeBackendLogsListEndpoint proxies a request to a worker node's /v1/backend-logs // endpoint to list model IDs that have backend logs. -func NodeBackendLogsListEndpoint(registry *nodes.NodeRegistry, registrationToken string) echo.HandlerFunc { +func NodeBackendLogsListEndpoint(registry *nodes.NodeRegistry, registrationToken string, dialFor nodes.WorkerNetDialerFor) echo.HandlerFunc { return func(c echo.Context) error { ctx := c.Request().Context() nodeID := c.Param("id") @@ -756,7 +756,7 @@ func NodeBackendLogsListEndpoint(registry *nodes.NodeRegistry, registrationToken return c.JSON(http.StatusBadGateway, nodeError(http.StatusBadGateway, "node has no HTTP address")) } - resp, err := proxyHTTPToWorker(node.HTTPAddress, "/v1/backend-logs", registrationToken) + resp, err := proxyHTTPToWorker(ctx, dialFor, node.ID, node.HTTPAddress, "/v1/backend-logs", registrationToken) if err != nil { return c.JSON(http.StatusBadGateway, nodeError(http.StatusBadGateway, fmt.Sprintf("failed to reach worker: %v", err))) } @@ -771,7 +771,7 @@ func NodeBackendLogsListEndpoint(registry *nodes.NodeRegistry, registrationToken // NodeBackendLogsLinesEndpoint proxies a request to a worker node's // /v1/backend-logs/{modelId} endpoint to get buffered log lines. -func NodeBackendLogsLinesEndpoint(registry *nodes.NodeRegistry, registrationToken string) echo.HandlerFunc { +func NodeBackendLogsLinesEndpoint(registry *nodes.NodeRegistry, registrationToken string, dialFor nodes.WorkerNetDialerFor) echo.HandlerFunc { return func(c echo.Context) error { ctx := c.Request().Context() nodeID := c.Param("id") @@ -787,7 +787,7 @@ func NodeBackendLogsLinesEndpoint(registry *nodes.NodeRegistry, registrationToke } path := "/v1/backend-logs/" + url.PathEscape(modelID) - resp, err := proxyHTTPToWorker(node.HTTPAddress, path, registrationToken) + resp, err := proxyHTTPToWorker(ctx, dialFor, node.ID, node.HTTPAddress, path, registrationToken) if err != nil { return c.JSON(http.StatusBadGateway, nodeError(http.StatusBadGateway, fmt.Sprintf("failed to reach worker: %v", err))) } @@ -802,7 +802,7 @@ func NodeBackendLogsLinesEndpoint(registry *nodes.NodeRegistry, registrationToke // NodeBackendLogsWSEndpoint proxies a WebSocket connection to a worker node's // /v1/backend-logs/{modelId}/ws endpoint for real-time log streaming. -func NodeBackendLogsWSEndpoint(registry *nodes.NodeRegistry, registrationToken string) echo.HandlerFunc { +func NodeBackendLogsWSEndpoint(registry *nodes.NodeRegistry, registrationToken string, dialFor nodes.WorkerNetDialerFor) echo.HandlerFunc { browserUpgrader := websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { origin := r.Header.Get("Origin") @@ -841,7 +841,7 @@ func NodeBackendLogsWSEndpoint(registry *nodes.NodeRegistry, registrationToken s workerHeaders.Set("Authorization", "Bearer "+registrationToken) } - workerDialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second} + workerDialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second, NetDialContext: dialFor(node.ID)} workerWS, _, err := workerDialer.Dial(workerURL, workerHeaders) if err != nil { browserWS.WriteMessage(websocket.CloseMessage, @@ -1312,9 +1312,14 @@ func DeleteSchedulingEndpoint(registry *nodes.NodeRegistry) echo.HandlerFunc { } // proxyHTTPToWorker makes a GET request to a worker's HTTP server with bearer token auth. -func proxyHTTPToWorker(httpAddress, path, token string) (*http.Response, error) { +// The connection goes through dialFor(nodeID) because the advertised address +// alone does not say how this frontend reaches that worker. +func proxyHTTPToWorker(ctx context.Context, dialFor nodes.WorkerNetDialerFor, nodeID, httpAddress, path, token string) (*http.Response, error) { reqURL := fmt.Sprintf("http://%s%s", httpAddress, path) - req, err := http.NewRequest("GET", reqURL, nil) + // WithoutCancel keeps the request bounded only by the 15s client timeout, + // as before the dialer change; cancelling on admin disconnect would be a + // separate, deliberate behaviour change. + req, err := http.NewRequestWithContext(context.WithoutCancel(ctx), "GET", reqURL, nil) if err != nil { return nil, err } @@ -1322,6 +1327,8 @@ func proxyHTTPToWorker(httpAddress, path, token string) (*http.Response, error) req.Header.Set("Authorization", "Bearer "+token) } - client := httpclient.NewWithTimeout(15 * time.Second) + t := httpclient.HardenedTransport() + t.DialContext = dialFor(nodeID) + client := httpclient.NewWithTimeout(15*time.Second, httpclient.WithTransport(t)) return client.Do(req) } diff --git a/core/http/endpoints/localai/nodes_backends_list_test.go b/core/http/endpoints/localai/nodes_backends_list_test.go index 636ab58b8..b9f0a0640 100644 --- a/core/http/endpoints/localai/nodes_backends_list_test.go +++ b/core/http/endpoints/localai/nodes_backends_list_test.go @@ -7,9 +7,9 @@ import ( "net/http/httptest" "github.com/labstack/echo/v4" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes" "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -21,21 +21,21 @@ type stubNodeCommandSender struct { listBackendsCalled bool } -func (s *stubNodeCommandSender) InstallBackend(_, _, _, _, _, _, _ string, _ int, _ string, _ func(messaging.BackendInstallProgressEvent)) (*messaging.BackendInstallReply, error) { - return &messaging.BackendInstallReply{}, nil +func (s *stubNodeCommandSender) InstallBackend(_, _, _, _, _, _, _ string, _ int, _ string, _ func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { + return &workerctl.BackendInstallReply{}, nil } -func (s *stubNodeCommandSender) UpgradeBackend(_, _, _, _, _, _ string, _ int, _ string, _ func(messaging.BackendInstallProgressEvent)) (*messaging.BackendUpgradeReply, error) { - return &messaging.BackendUpgradeReply{}, nil +func (s *stubNodeCommandSender) UpgradeBackend(_, _, _, _, _, _ string, _ int, _ string, _ func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendUpgradeReply, error) { + return &workerctl.BackendUpgradeReply{}, nil } -func (s *stubNodeCommandSender) DeleteBackend(_, _ string) (*messaging.BackendDeleteReply, error) { - return &messaging.BackendDeleteReply{Success: true}, nil +func (s *stubNodeCommandSender) DeleteBackend(_, _ string) (*workerctl.BackendDeleteReply, error) { + return &workerctl.BackendDeleteReply{Success: true}, nil } -func (s *stubNodeCommandSender) ListBackends(_ string) (*messaging.BackendListReply, error) { +func (s *stubNodeCommandSender) ListBackends(_ string) (*workerctl.BackendListReply, error) { s.listBackendsCalled = true - return &messaging.BackendListReply{Backends: []messaging.NodeBackendInfo{{Name: "llama-cpp"}}}, nil + return &workerctl.BackendListReply{Backends: []workerctl.NodeBackendInfo{{Name: "llama-cpp"}}}, nil } func (s *stubNodeCommandSender) StopBackend(_, _ string) error { return nil } @@ -78,7 +78,7 @@ var _ = Describe("ListBackendsOnNodeEndpoint", func() { Expect(stub.listBackendsCalled).To(BeFalse(), "agent workers don't subscribe to backend.list; the endpoint must not issue the doomed NATS request") - var list []messaging.NodeBackendInfo + var list []workerctl.NodeBackendInfo Expect(json.Unmarshal(rec.Body.Bytes(), &list)).To(Succeed()) Expect(list).To(BeEmpty()) // Must be `[]`, not `null`, so the UI can render it. @@ -97,7 +97,7 @@ var _ = Describe("ListBackendsOnNodeEndpoint", func() { Expect(stub.listBackendsCalled).To(BeTrue(), "backend nodes must still be queried over NATS") - var list []messaging.NodeBackendInfo + var list []workerctl.NodeBackendInfo Expect(json.Unmarshal(rec.Body.Bytes(), &list)).To(Succeed()) Expect(list).To(HaveLen(1)) Expect(list[0].Name).To(Equal("llama-cpp")) diff --git a/core/http/endpoints/mcp/executor.go b/core/http/endpoints/mcp/executor.go index 9f9b279d6..037ae0813 100644 --- a/core/http/endpoints/mcp/executor.go +++ b/core/http/endpoints/mcp/executor.go @@ -58,28 +58,28 @@ func (e *LocalToolExecutor) HasTools() bool { return len(e.tools) > 0 } -// DistributedToolExecutor routes tool operations through NATS to agent workers. +// DistributedToolExecutor routes tool operations to agent workers. type DistributedToolExecutor struct { - natsClient MCPNATSClient - modelName string - remote config.MCPGenericConfig[config.MCPRemoteServers] - stdio config.MCPGenericConfig[config.MCPSTDIOServers] - toolDefs []mcpRemote.MCPToolDef + agentControl AgentControl + modelName string + remote config.MCPGenericConfig[config.MCPRemoteServers] + stdio config.MCPGenericConfig[config.MCPSTDIOServers] + toolDefs []mcpRemote.MCPToolDef } -// NewDistributedToolExecutor creates a ToolExecutor that routes through NATS. -// It discovers tools immediately via a NATS request-reply to an agent worker. -func NewDistributedToolExecutor(ctx context.Context, natsClient MCPNATSClient, modelName string, +// NewDistributedToolExecutor creates a ToolExecutor that routes to agent workers. +// It discovers tools immediately with a discovery request to an agent worker. +func NewDistributedToolExecutor(ctx context.Context, agentControl AgentControl, modelName string, remote config.MCPGenericConfig[config.MCPRemoteServers], stdio config.MCPGenericConfig[config.MCPSTDIOServers], ) *DistributedToolExecutor { e := &DistributedToolExecutor{ - natsClient: natsClient, - modelName: modelName, - remote: remote, - stdio: stdio, + agentControl: agentControl, + modelName: modelName, + remote: remote, + stdio: stdio, } - resp, err := DiscoverMCPToolsRemote(ctx, natsClient, modelName, remote, stdio) + resp, err := DiscoverMCPToolsRemote(ctx, agentControl, modelName, remote, stdio) if err != nil { xlog.Error("Failed to discover MCP tools (distributed)", "error", err) } else if resp != nil { @@ -103,7 +103,7 @@ func (e *DistributedToolExecutor) IsTool(name string) bool { } func (e *DistributedToolExecutor) ExecuteTool(ctx context.Context, toolName, arguments string) (string, error) { - return ExecuteMCPToolCallRemote(ctx, e.natsClient, e.modelName, e.remote, e.stdio, toolName, arguments) + return ExecuteMCPToolCallRemote(ctx, e.agentControl, e.modelName, e.remote, e.stdio, toolName, arguments) } func (e *DistributedToolExecutor) HasTools() bool { @@ -111,15 +111,15 @@ func (e *DistributedToolExecutor) HasTools() bool { } // NewToolExecutor creates the appropriate ToolExecutor based on the current mode. -// When natsClient is non-nil, returns a DistributedToolExecutor that routes through NATS. -// When natsClient is nil, creates local sessions and returns a LocalToolExecutor. -func NewToolExecutor(ctx context.Context, natsClient MCPNATSClient, modelName string, +// When agentControl is non-nil, returns a DistributedToolExecutor that routes to agent workers. +// When agentControl is nil, creates local sessions and returns a LocalToolExecutor. +func NewToolExecutor(ctx context.Context, agentControl AgentControl, modelName string, remote config.MCPGenericConfig[config.MCPRemoteServers], stdio config.MCPGenericConfig[config.MCPSTDIOServers], enabledServers []string, ) ToolExecutor { - if natsClient != nil { - return NewDistributedToolExecutor(ctx, natsClient, modelName, remote, stdio) + if agentControl != nil { + return NewDistributedToolExecutor(ctx, agentControl, modelName, remote, stdio) } sessions, err := NamedSessionsFromMCPConfig(modelName, remote, stdio, enabledServers) if err != nil || len(sessions) == 0 { diff --git a/core/http/endpoints/mcp/executor_agent_control_test.go b/core/http/endpoints/mcp/executor_agent_control_test.go new file mode 100644 index 000000000..15423f76a --- /dev/null +++ b/core/http/endpoints/mcp/executor_agent_control_test.go @@ -0,0 +1,76 @@ +package mcp + +import ( + "context" + "time" + + "github.com/mudler/LocalAI/core/config" + mcpRemote "github.com/mudler/LocalAI/core/services/mcp" + "github.com/mudler/LocalAI/pkg/functions" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type recordingAgentControl struct { + toolReply mcpRemote.MCPToolResponse + discoveryReply mcpRemote.MCPDiscoveryResponse + discoveries int + toolDeadline time.Time + discDeadline time.Time +} + +func (r *recordingAgentControl) ExecuteMCPTool(ctx context.Context, _ mcpRemote.MCPToolRequest) (*mcpRemote.MCPToolResponse, error) { + r.toolDeadline, _ = ctx.Deadline() + reply := r.toolReply + return &reply, nil +} + +func (r *recordingAgentControl) DiscoverMCPTools(ctx context.Context, _ mcpRemote.MCPDiscoveryRequest) (*mcpRemote.MCPDiscoveryResponse, error) { + r.discoveries++ + r.discDeadline, _ = ctx.Deadline() + reply := r.discoveryReply + return &reply, nil +} + +// The distributed mode switch is interface nil-ness: a nil AgentControl must +// keep MCP sessions local, exactly as a nil messaging client did before. +var _ = Describe("MCP routing through AgentControl", func() { + var ( + remote config.MCPGenericConfig[config.MCPRemoteServers] + stdio config.MCPGenericConfig[config.MCPSTDIOServers] + ) + + It("keeps sessions local when no agent control is wired", func() { + exec := NewToolExecutor(context.Background(), nil, "agent-control-nil", remote, stdio, nil) + Expect(exec).To(BeAssignableToTypeOf(&LocalToolExecutor{})) + }) + + It("routes to agent workers when agent control is wired", func() { + ac := &recordingAgentControl{discoveryReply: mcpRemote.MCPDiscoveryResponse{ + Tools: []mcpRemote.MCPToolDef{{ToolName: "weather", Function: functions.Function{Name: "weather"}}}, + }} + exec := NewToolExecutor(context.Background(), ac, "agent-control-set", remote, stdio, nil) + Expect(exec).To(BeAssignableToTypeOf(&DistributedToolExecutor{})) + Expect(ac.discoveries).To(Equal(1)) + Expect(exec.IsTool("weather")).To(BeTrue()) + }) + + It("keeps the worker's tool error text and bounds the call by the tool budget", func() { + ac := &recordingAgentControl{toolReply: mcpRemote.MCPToolResponse{Error: "tool 'x' not found"}} + start := time.Now() + _, err := ExecuteMCPToolCallRemote(context.Background(), ac, "m", remote, stdio, "x", "{}") + Expect(err).To(MatchError("remote MCP tool error: tool 'x' not found")) + Expect(ac.toolDeadline).ToNot(BeZero()) + Expect(ac.toolDeadline.Sub(start)).To(BeNumerically("~", config.DefaultMCPToolTimeout, time.Second)) + }) + + It("keeps the worker's discovery error text and bounds the call by the discovery budget", func() { + ac := &recordingAgentControl{discoveryReply: mcpRemote.MCPDiscoveryResponse{Error: "no MCP servers"}} + start := time.Now() + _, err := DiscoverMCPToolsRemote(context.Background(), ac, "m", remote, stdio) + Expect(err).To(MatchError("remote MCP discovery error: no MCP servers")) + Expect(ac.discDeadline).ToNot(BeZero()) + Expect(ac.discDeadline.Sub(start)).To(BeNumerically("~", config.DefaultMCPDiscoveryTimeout, time.Second)) + }) +}) diff --git a/core/http/endpoints/mcp/tools.go b/core/http/endpoints/mcp/tools.go index 0b4931a05..02ea46037 100644 --- a/core/http/endpoints/mcp/tools.go +++ b/core/http/endpoints/mcp/tools.go @@ -16,7 +16,6 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/schema" mcpRemote "github.com/mudler/LocalAI/core/services/mcp" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/pkg/functions" "github.com/mudler/LocalAI/pkg/httpclient" @@ -100,9 +99,14 @@ var ( client = mcp.NewClient(&mcp.Implementation{Name: "LocalAI", Version: "v1.0.0"}, nil) ) -// MCPNATSClient is the interface for NATS request-reply operations needed by MCP routing. -type MCPNATSClient interface { - Request(subject string, data []byte, timeout time.Duration) ([]byte, error) +// AgentControl carries the frontend's MCP verbs to one agent worker. A decoded +// reply is returned with a nil error even when its Error field is set: that is +// the worker's own answer. nodes.ErrNoRoute (wrapped) means no agent worker +// could be offered the request. A timeout, a transport fault or an unreadable +// reply is an ordinary error and never ErrNoRoute. +type AgentControl interface { + ExecuteMCPTool(ctx context.Context, req mcpRemote.MCPToolRequest) (*mcpRemote.MCPToolResponse, error) + DiscoverMCPTools(ctx context.Context, req mcpRemote.MCPDiscoveryRequest) (*mcpRemote.MCPDiscoveryResponse, error) } // MetadataKeyLocalAIAssistant is the request-metadata key the chat handler @@ -510,18 +514,18 @@ func ExecuteMCPToolCall(ctx context.Context, tools []MCPToolInfo, toolName strin return string(combined), nil } -// ExecuteMCPToolCallRemote routes an MCP tool execution request to an agent worker via NATS. +// ExecuteMCPToolCallRemote routes an MCP tool execution request to an agent worker. // Used in distributed mode when the frontend doesn't hold MCP sessions locally. func ExecuteMCPToolCallRemote( ctx context.Context, - natsClient MCPNATSClient, + agentControl AgentControl, modelName string, remote config.MCPGenericConfig[config.MCPRemoteServers], stdio config.MCPGenericConfig[config.MCPSTDIOServers], toolName, arguments string, ) (string, error) { - if natsClient == nil { - return "", fmt.Errorf("NATS client not configured for distributed MCP") + if agentControl == nil { + return "", fmt.Errorf("agent control not configured for distributed MCP") } var args map[string]any @@ -538,16 +542,12 @@ func ExecuteMCPToolCallRemote( RemoteServers: remote, StdioServers: stdio, } - reqData, _ := json.Marshal(req) - replyData, err := natsClient.Request(messaging.SubjectMCPToolExecute, reqData, config.DefaultMCPToolTimeout) + ctx, cancel := context.WithTimeout(ctx, config.DefaultMCPToolTimeout) + defer cancel() + resp, err := agentControl.ExecuteMCPTool(ctx, req) if err != nil { - return "", fmt.Errorf("NATS MCP tool request failed: %w", err) - } - - var resp mcpRemote.MCPToolResponse - if err := json.Unmarshal(replyData, &resp); err != nil { - return "", fmt.Errorf("unmarshal MCP reply: %w", err) + return "", fmt.Errorf("remote MCP tool request failed: %w", err) } if resp.Error != "" { return "", fmt.Errorf("remote MCP tool error: %s", resp.Error) @@ -555,17 +555,17 @@ func ExecuteMCPToolCallRemote( return resp.Result, nil } -// DiscoverMCPToolsRemote routes an MCP discovery request to an agent worker via NATS. +// DiscoverMCPToolsRemote routes an MCP discovery request to an agent worker. // Returns server info and tool function schemas from the remote worker. func DiscoverMCPToolsRemote( ctx context.Context, - natsClient MCPNATSClient, + agentControl AgentControl, modelName string, remote config.MCPGenericConfig[config.MCPRemoteServers], stdio config.MCPGenericConfig[config.MCPSTDIOServers], ) (*mcpRemote.MCPDiscoveryResponse, error) { - if natsClient == nil { - return nil, fmt.Errorf("NATS client not configured for distributed MCP") + if agentControl == nil { + return nil, fmt.Errorf("agent control not configured for distributed MCP") } req := mcpRemote.MCPDiscoveryRequest{ @@ -573,21 +573,17 @@ func DiscoverMCPToolsRemote( RemoteServers: remote, StdioServers: stdio, } - reqData, _ := json.Marshal(req) - replyData, err := natsClient.Request(messaging.SubjectMCPDiscovery, reqData, config.DefaultMCPDiscoveryTimeout) + ctx, cancel := context.WithTimeout(ctx, config.DefaultMCPDiscoveryTimeout) + defer cancel() + resp, err := agentControl.DiscoverMCPTools(ctx, req) if err != nil { - return nil, fmt.Errorf("NATS MCP discovery request failed: %w", err) - } - - var resp mcpRemote.MCPDiscoveryResponse - if err := json.Unmarshal(replyData, &resp); err != nil { - return nil, fmt.Errorf("unmarshal MCP discovery reply: %w", err) + return nil, fmt.Errorf("remote MCP discovery request failed: %w", err) } if resp.Error != "" { return nil, fmt.Errorf("remote MCP discovery error: %s", resp.Error) } - return &resp, nil + return resp, nil } // ListMCPServers returns server info with tool, prompt, and resource names for each session. diff --git a/core/http/endpoints/openai/chat.go b/core/http/endpoints/openai/chat.go index 76efbca28..638be16a9 100644 --- a/core/http/endpoints/openai/chat.go +++ b/core/http/endpoints/openai/chat.go @@ -218,7 +218,7 @@ func applyAutoparserOverride( // @Param request body schema.OpenAIRequest true "query params" // @Success 200 {object} schema.OpenAIResponse "Response" // @Router /v1/chat/completions [post] -func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, startupOptions *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient, assistantHolder *mcpTools.LocalAIAssistantHolder, compressor middleware.ChatCompressor) echo.HandlerFunc { +func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, startupOptions *config.ApplicationConfig, agentControl mcpTools.AgentControl, assistantHolder *mcpTools.LocalAIAssistantHolder, compressor middleware.ChatCompressor) echo.HandlerFunc { return func(c echo.Context) error { var textContentToReturn string id := uuid.New().String() @@ -320,7 +320,7 @@ func ChatEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator if (len(mcpServers) > 0 || mcpPromptName != "" || len(mcpResourceURIs) > 0) && (config.MCP.Servers != "" || config.MCP.Stdio != "") { remote, stdio, mcpErr := config.MCP.MCPConfigFromYAML() if mcpErr == nil { - mcpExecutor = mcpTools.NewToolExecutor(c.Request().Context(), natsClient, config.Name, remote, stdio, mcpServers) + mcpExecutor = mcpTools.NewToolExecutor(c.Request().Context(), agentControl, config.Name, remote, stdio, mcpServers) // Prompt and resource injection (pre-processing step — resolves locally regardless of distributed mode) namedSessions, sessErr := mcpTools.NamedSessionsFromMCPConfig(config.Name, remote, stdio, mcpServers) diff --git a/core/http/endpoints/openresponses/responses.go b/core/http/endpoints/openresponses/responses.go index 528737273..71164e291 100644 --- a/core/http/endpoints/openresponses/responses.go +++ b/core/http/endpoints/openresponses/responses.go @@ -31,7 +31,7 @@ import ( // @Param request body schema.OpenResponsesRequest true "Request body" // @Success 200 {object} schema.ORResponseResource "Response" // @Router /v1/responses [post] -func ResponsesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, natsClient mcpTools.MCPNATSClient) echo.HandlerFunc { +func ResponsesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, evaluator *templates.Evaluator, appConfig *config.ApplicationConfig, agentControl mcpTools.AgentControl) echo.HandlerFunc { return func(c echo.Context) error { createdAt := time.Now().Unix() responseID := fmt.Sprintf("resp_%s", uuid.New().String()) @@ -108,7 +108,7 @@ func ResponsesEndpoint(cl *config.ModelConfigLoader, ml *model.ModelLoader, eval if !hasMCPRequest { enabledServers = nil // backward compat: auto-activate all servers } - mcpExecutor = mcpTools.NewToolExecutor(c.Request().Context(), natsClient, cfg.Name, remote, stdio, enabledServers) + mcpExecutor = mcpTools.NewToolExecutor(c.Request().Context(), agentControl, cfg.Name, remote, stdio, enabledServers) // Prompt and resource injection (pre-processing step — resolves locally regardless of distributed mode) if hasMCPRequest { diff --git a/core/http/endpoints/openresponses/store.go b/core/http/endpoints/openresponses/store.go index f703e54b1..79ac99642 100644 --- a/core/http/endpoints/openresponses/store.go +++ b/core/http/endpoints/openresponses/store.go @@ -56,7 +56,7 @@ type ResponseStore struct { // (see sync.go), which is how a standalone deployment keeps exactly the // previous process-local behaviour. Guarded by mu. synced *syncstate.SyncedMap[string, *syncedResponse] - nats messaging.MessagingClient + nats messaging.Broadcaster cancelSub messaging.Subscription replicaID string lifeCtx context.Context diff --git a/core/http/endpoints/openresponses/sync.go b/core/http/endpoints/openresponses/sync.go index bd15b39cd..3240d2cb4 100644 --- a/core/http/endpoints/openresponses/sync.go +++ b/core/http/endpoints/openresponses/sync.go @@ -80,7 +80,7 @@ type responseCancelEvent struct { // through deltas alone. A replica that joins later does not learn about // responses created before it started; that is the same visibility a client had // before this change and strictly better than the 404 it got from every peer. -func (s *ResponseStore) EnableDistributed(ctx context.Context, nats messaging.MessagingClient, replicaID string) error { +func (s *ResponseStore) EnableDistributed(ctx context.Context, nats messaging.Broadcaster, replicaID string) error { if nats == nil { return nil } @@ -162,7 +162,7 @@ func (s *ResponseStore) syncMap() *syncstate.SyncedMap[string, *syncedResponse] // distributed returns the replication handles as a consistent snapshot. Every // path that broadcasts reads them through here so a concurrent Close cannot be // observed half-applied. A nil map means standalone mode. -func (s *ResponseStore) distributed() (*syncstate.SyncedMap[string, *syncedResponse], context.Context, messaging.MessagingClient, string) { +func (s *ResponseStore) distributed() (*syncstate.SyncedMap[string, *syncedResponse], context.Context, messaging.Broadcaster, string) { s.mu.RLock() defer s.mu.RUnlock() ctx := s.lifeCtx diff --git a/core/http/routes/anthropic.go b/core/http/routes/anthropic.go index 124557655..a9e6a48ef 100644 --- a/core/http/routes/anthropic.go +++ b/core/http/routes/anthropic.go @@ -25,9 +25,9 @@ func RegisterAnthropicRoutes(app *echo.Echo, application *application.Application, ) { // Anthropic Messages API endpoint - var natsClient mcpTools.MCPNATSClient + var agentControl mcpTools.AgentControl if d := application.Distributed(); d != nil { - natsClient = d.Nats + agentControl = d.AgentControl } messagesHandler := anthropic.MessagesEndpoint( @@ -35,7 +35,7 @@ func RegisterAnthropicRoutes(app *echo.Echo, application.ModelLoader(), application.TemplatesEvaluator(), application.ApplicationConfig(), - natsClient, + agentControl, ) messagesMiddleware := []echo.MiddlewareFunc{ diff --git a/core/http/routes/localai.go b/core/http/routes/localai.go index 7a6f901ec..b2bfed146 100644 --- a/core/http/routes/localai.go +++ b/core/http/routes/localai.go @@ -476,11 +476,11 @@ func RegisterLocalAIRoutes(router *echo.Echo, compressionservice.CounterFunc(tokens.CountMessages), compressionservice.NewInferenceSummarizer(cl, ml, appConfig), ) - var mcpNATS mcpTools.MCPNATSClient + var agentControl mcpTools.AgentControl if d := app.Distributed(); d != nil { - mcpNATS = d.Nats + agentControl = d.AgentControl } - mcpStreamHandler := localai.MCPEndpoint(cl, ml, evaluator, appConfig, mcpNATS, chatCompressor) + mcpStreamHandler := localai.MCPEndpoint(cl, ml, evaluator, appConfig, agentControl, chatCompressor) mcpStreamMiddleware := []echo.MiddlewareFunc{ requestExtractor.BuildFilteredFirstAvailableDefaultModel(config.BuildUsecaseFilterFn(config.FLAG_CHAT)), requestExtractor.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }), @@ -499,7 +499,7 @@ func RegisterLocalAIRoutes(router *echo.Echo, router.POST("/mcp/chat/completions", mcpStreamHandler, mcpStreamMiddleware...) // MCP server listing endpoint - router.GET("/v1/mcp/servers/:model", localai.MCPServersEndpoint(cl, appConfig, mcpNATS), mcpMw) + router.GET("/v1/mcp/servers/:model", localai.MCPServersEndpoint(cl, appConfig, agentControl), mcpMw) // MCP prompts endpoints router.GET("/v1/mcp/prompts/:model", localai.MCPPromptsEndpoint(cl, appConfig), mcpMw) diff --git a/core/http/routes/nodes.go b/core/http/routes/nodes.go index 053d6c19c..eeb2819e9 100644 --- a/core/http/routes/nodes.go +++ b/core/http/routes/nodes.go @@ -61,7 +61,7 @@ func RegisterNodeSelfServiceRoutes(e *echo.Echo, registry *nodes.NodeRegistry, r // backend install path (POST /:id/backends/install). That handler enqueues a // ManagementOp on the gallery channel rather than blocking on a NATS reply, so // the browser gets HTTP 202 + jobID immediately instead of waiting up to 3 minutes. -func RegisterNodeAdminRoutes(e *echo.Echo, registry *nodes.NodeRegistry, unloader nodes.NodeCommandSender, galleryService *galleryop.GalleryService, opcache *galleryop.OpCache, appConfig *config.ApplicationConfig, adminMw echo.MiddlewareFunc, authDB *gorm.DB, hmacSecret string, registrationToken string, natsCfg natsauth.Config) { +func RegisterNodeAdminRoutes(e *echo.Echo, registry *nodes.NodeRegistry, unloader nodes.NodeCommandSender, galleryService *galleryop.GalleryService, opcache *galleryop.OpCache, appConfig *config.ApplicationConfig, adminMw echo.MiddlewareFunc, authDB *gorm.DB, hmacSecret string, registrationToken string, natsCfg natsauth.Config, workerHTTPDial nodes.WorkerNetDialerFor) { if registry == nil { return } @@ -101,8 +101,8 @@ func RegisterNodeAdminRoutes(e *echo.Echo, registry *nodes.NodeRegistry, unloade admin.POST("/:id/models/delete", localai.DeleteModelOnNodeEndpoint(unloader, registry)) // Backend log streaming (proxied from worker HTTP server) - admin.GET("/:id/backend-logs", localai.NodeBackendLogsListEndpoint(registry, registrationToken)) - admin.GET("/:id/backend-logs/:modelId", localai.NodeBackendLogsLinesEndpoint(registry, registrationToken)) + admin.GET("/:id/backend-logs", localai.NodeBackendLogsListEndpoint(registry, registrationToken, workerHTTPDial)) + admin.GET("/:id/backend-logs/:modelId", localai.NodeBackendLogsLinesEndpoint(registry, registrationToken, workerHTTPDial)) // Label management admin.GET("/:id/labels", localai.GetNodeLabelsEndpoint(registry)) @@ -123,7 +123,7 @@ func RegisterNodeAdminRoutes(e *echo.Echo, registry *nodes.NodeRegistry, unloade admin.DELETE("/:id/vram-budget", localai.ResetVRAMBudgetEndpoint(registry)) // WebSocket proxy for real-time log streaming from workers - e.GET("/ws/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsWSEndpoint(registry, registrationToken), readyMw, adminMw) + e.GET("/ws/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsWSEndpoint(registry, registrationToken, workerHTTPDial), readyMw, adminMw) } // nodeTokenAuth validates the registration token for node self-service endpoints. diff --git a/core/http/routes/openai.go b/core/http/routes/openai.go index eb752693d..ab9ceb3fa 100644 --- a/core/http/routes/openai.go +++ b/core/http/routes/openai.go @@ -38,10 +38,10 @@ func RegisterOpenAIRoutes(app *echo.Echo, app.POST("/v1/realtime/transcription_session", openai.RealtimeTranscriptionSession(application), traceMiddleware) app.POST("/v1/realtime/calls", openai.RealtimeCalls(application), traceMiddleware) - // NATS client for distributed MCP tool routing (nil when not in distributed mode) - var natsClient mcpTools.MCPNATSClient + // Agent control for distributed MCP tool routing (nil when not in distributed mode) + var agentControl mcpTools.AgentControl if d := application.Distributed(); d != nil { - natsClient = d.Nats + agentControl = d.AgentControl } // chat @@ -49,7 +49,7 @@ func RegisterOpenAIRoutes(app *echo.Echo, compressionservice.CounterFunc(tokens.CountMessages), compressionservice.NewInferenceSummarizer(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig()), ) - chatHandler := openai.ChatEndpoint(application.ModelConfigLoader(), application.ModelLoader(), application.TemplatesEvaluator(), application.ApplicationConfig(), natsClient, application.LocalAIAssistant(), chatCompressor) + chatHandler := openai.ChatEndpoint(application.ModelConfigLoader(), application.ModelLoader(), application.TemplatesEvaluator(), application.ApplicationConfig(), agentControl, application.LocalAIAssistant(), chatCompressor) chatMiddleware := []echo.MiddlewareFunc{ nodeHeaderMiddleware, usageMiddleware, diff --git a/core/http/routes/openresponses.go b/core/http/routes/openresponses.go index 8aff6ccf1..101567a6a 100644 --- a/core/http/routes/openresponses.go +++ b/core/http/routes/openresponses.go @@ -16,10 +16,10 @@ func RegisterOpenResponsesRoutes(app *echo.Echo, re *middleware.RequestExtractor, application *application.Application) { - // NATS client for distributed MCP tool routing (nil when not in distributed mode) - var natsClient mcpTools.MCPNATSClient + // Agent control for distributed MCP tool routing (nil when not in distributed mode) + var agentControl mcpTools.AgentControl if d := application.Distributed(); d != nil { - natsClient = d.Nats + agentControl = d.AgentControl // Replicate response metadata across frontend replicas and subscribe to // delegated cancels. Without this a GET, a previous_response_id lookup or @@ -38,7 +38,7 @@ func RegisterOpenResponsesRoutes(app *echo.Echo, application.ModelLoader(), application.TemplatesEvaluator(), application.ApplicationConfig(), - natsClient, + agentControl, ) responsesMiddleware := []echo.MiddlewareFunc{ diff --git a/core/services/agentpool/agent_jobs.go b/core/services/agentpool/agent_jobs.go index 59850981a..0fcf47c49 100644 --- a/core/services/agentpool/agent_jobs.go +++ b/core/services/agentpool/agent_jobs.go @@ -54,7 +54,7 @@ type AgentJobService struct { // taskNats is the distributed NATS client backing the tasks SyncedMap. It is // not available at construction time, so it is injected via SetTaskSyncNATS // during distributed wiring; nil keeps tasks in-memory-only (standalone). - taskNats messaging.MessagingClient + taskNats messaging.Broadcaster // Storage (in-memory primary, persister for secondary persistence) jobs *xsync.SyncedMap[string, schema.Job] @@ -115,7 +115,7 @@ func (s *AgentJobService) SetDistributedJobStore(store *jobs.JobStore) { // tasks SyncedMap is rebuilt to pick it up. It is always called before Start / // hydrate, while the map is still empty, so rebuilding loses no state. Passing nil // (standalone) keeps the map in-memory-only with no broadcast. -func (s *AgentJobService) SetTaskSyncNATS(nats messaging.MessagingClient) { +func (s *AgentJobService) SetTaskSyncNATS(nats messaging.Broadcaster) { s.taskNats = nats s.buildTasksMap() } diff --git a/core/services/agentpool/agent_pool.go b/core/services/agentpool/agent_pool.go index dd9c49da8..a4fa06a88 100644 --- a/core/services/agentpool/agent_pool.go +++ b/core/services/agentpool/agent_pool.go @@ -69,11 +69,10 @@ type localAGICore struct { // distributedBridge connects to the NATS-based distributed agent system. type distributedBridge struct { - natsClient messaging.Publisher // NATS client for distributed agent execution + workQueue messaging.WorkQueue // Non-nil selects distributed mode; carries agent runs agentStore *agents.AgentStore // PostgreSQL agent config store eventBridge AgentEventBridge // Event bridge for SSE + persistence skillStore *distributed.SkillStore // PostgreSQL skill metadata (distributed mode) - dispatcher agents.Dispatcher // Native dispatcher (distributed or local) } // userManager handles per-user services, storage, and auth. @@ -123,7 +122,7 @@ type AgentConfigStore interface { type AgentPoolOptions struct { AuthDB *gorm.DB SkillStore *distributed.SkillStore - NATSClient messaging.Publisher + WorkQueue messaging.WorkQueue EventBridge AgentEventBridge AgentStore *agents.AgentStore } @@ -140,8 +139,8 @@ func NewAgentPoolService(appConfig *config.ApplicationConfig, opts ...AgentPoolO if o.SkillStore != nil { svc.distributed.skillStore = o.SkillStore } - if o.NATSClient != nil { - svc.distributed.natsClient = o.NATSClient + if o.WorkQueue != nil { + svc.distributed.workQueue = o.WorkQueue } if o.EventBridge != nil { svc.distributed.eventBridge = o.EventBridge @@ -175,7 +174,7 @@ func (s *AgentPoolService) Start(ctx context.Context) error { // Distributed mode: use native executor + NATSDispatcher. // No LocalAGI pool, no collections, no skills service — all stateless. - if s.distributed.natsClient != nil { + if s.distributed.workQueue != nil { return s.startDistributed(ctx, apiURL, apiKey) } @@ -244,16 +243,15 @@ func (s *AgentPoolService) startDistributed(ctx context.Context, apiURL, apiKey // Start the background agent scheduler on the frontend. // It needs DB access to list configs and update LastRunAt — the worker doesn't have DB. // The advisory lock ensures only one frontend instance runs the scheduler. - if s.users.authDB != nil && s.distributed.natsClient != nil && s.distributed.agentStore != nil { + if s.users.authDB != nil && s.distributed.workQueue != nil && s.distributed.agentStore != nil { var schedulerOpts []agents.AgentSchedulerOpt if s.distributed.skillStore != nil { schedulerOpts = append(schedulerOpts, agents.WithSchedulerSkillProvider(s.buildSkillProvider())) } scheduler := agents.NewAgentScheduler( s.users.authDB, - s.distributed.natsClient, + s.distributed.workQueue, s.distributed.agentStore, - messaging.SubjectAgentExecute, schedulerOpts..., ) go scheduler.Start(ctx) @@ -391,12 +389,6 @@ func (s *AgentPoolService) Pool() *state.AgentPool { return s.localAGI.pool } -// SetNATSClient sets the NATS client for distributed agent execution. -// Deprecated: prefer passing NATSClient via AgentPoolOptions at construction time. -func (s *AgentPoolService) SetNATSClient(nc messaging.Publisher) { - s.distributed.natsClient = nc -} - // SetEventBridge sets the event bridge for distributed SSE + persistence. // Deprecated: prefer passing EventBridge via AgentPoolOptions at construction time. func (s *AgentPoolService) SetEventBridge(eb AgentEventBridge) { @@ -996,7 +988,7 @@ func (s *AgentPoolService) ChatForUser(userID, name, message string) (string, er return s.configBackend.Chat(userID, name, message) } -// dispatchChat publishes a chat event to the NATS agent execution queue. +// dispatchChat enqueues a chat event as agent-run work. // The event is enriched with the full agent config and resolved skills so that // the worker does not need direct database access. func (s *AgentPoolService) dispatchChat(userID, name, message string) (string, error) { @@ -1040,7 +1032,7 @@ func (s *AgentPoolService) dispatchChat(userID, name, message string) (string, e Config: cfg, Skills: skills, } - if err := s.distributed.natsClient.Publish(messaging.SubjectAgentExecute, evt); err != nil { + if err := s.distributed.workQueue.Enqueue(context.Background(), messaging.WorkAgentRun, evt); err != nil { return "", fmt.Errorf("failed to dispatch agent chat: %w", err) } return messageID, nil diff --git a/core/services/agentpool/dispatch_chat_test.go b/core/services/agentpool/dispatch_chat_test.go new file mode 100644 index 000000000..1dc541c8b --- /dev/null +++ b/core/services/agentpool/dispatch_chat_test.go @@ -0,0 +1,65 @@ +package agentpool + +import ( + "context" + "errors" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/agents" + "github.com/mudler/LocalAI/core/services/messaging" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// recordingWorkQueue keeps the typed payload so a spec can tell an +// AgentChatEvent from pre-encoded bytes. +type recordingWorkQueue struct { + kinds []messaging.WorkKind + payloads []any + err error +} + +func (q *recordingWorkQueue) Enqueue(_ context.Context, kind messaging.WorkKind, payload any) error { + q.kinds = append(q.kinds, kind) + q.payloads = append(q.payloads, payload) + return q.err +} + +var _ = Describe("dispatchChat", func() { + It("enqueues an agent run carrying the typed chat event", func() { + queue := &recordingWorkQueue{} + svc, err := NewAgentPoolService(&config.ApplicationConfig{}, AgentPoolOptions{WorkQueue: queue}) + Expect(err).ToNot(HaveOccurred()) + + id, err := svc.dispatchChat("user-1", "agent-a", "hello") + Expect(err).ToNot(HaveOccurred()) + Expect(id).ToNot(BeEmpty()) + + Expect(queue.kinds).To(Equal([]messaging.WorkKind{messaging.WorkAgentRun})) + evt, ok := queue.payloads[0].(agents.AgentChatEvent) + Expect(ok).To(BeTrue(), "the queue must be handed the typed AgentChatEvent") + Expect(evt.AgentName).To(Equal("agent-a")) + Expect(evt.UserID).To(Equal("user-1")) + Expect(evt.Message).To(Equal("hello")) + Expect(evt.MessageID).To(Equal(id)) + Expect(evt.Role).To(Equal("user")) + }) + + It("reports an enqueue failure to the caller", func() { + queue := &recordingWorkQueue{err: errors.New("carrier down")} + svc, err := NewAgentPoolService(&config.ApplicationConfig{}, AgentPoolOptions{WorkQueue: queue}) + Expect(err).ToNot(HaveOccurred()) + + _, err = svc.dispatchChat("user-1", "agent-a", "hello") + Expect(err).To(MatchError(ContainSubstring("carrier down"))) + }) + + It("stays in local mode when no work queue is given", func() { + svc, err := NewAgentPoolService(&config.ApplicationConfig{}, AgentPoolOptions{}) + Expect(err).ToNot(HaveOccurred()) + // A strict interface comparison: Gomega's BeNil would also accept a + // typed-nil pointer, which the mode switch reads as distributed. + Expect(svc.distributed.workQueue == nil).To(BeTrue()) + }) +}) diff --git a/core/services/agentpool/user_services.go b/core/services/agentpool/user_services.go index 56d19e0fc..4c665994c 100644 --- a/core/services/agentpool/user_services.go +++ b/core/services/agentpool/user_services.go @@ -31,7 +31,7 @@ type UserServicesManager struct { jobDBStore *jobs.JobStore // jobNats keeps per-user agent tasks consistent across replicas (nil in // standalone). Inherited by each per-user AgentJobService. - jobNats messaging.MessagingClient + jobNats messaging.Broadcaster } // NewUserServicesManager creates a new UserServicesManager. @@ -199,7 +199,7 @@ func (m *UserServicesManager) SetJobDBStore(s *jobs.JobStore) { // SetJobSyncNATS sets the NATS client used to keep per-user agent tasks consistent // across replicas. -func (m *UserServicesManager) SetJobSyncNATS(nats messaging.MessagingClient) { +func (m *UserServicesManager) SetJobSyncNATS(nats messaging.Broadcaster) { m.jobNats = nats } diff --git a/core/services/agents/dispatcher.go b/core/services/agents/dispatcher.go index 3ed737e83..562be1db2 100644 --- a/core/services/agents/dispatcher.go +++ b/core/services/agents/dispatcher.go @@ -40,17 +40,6 @@ type AgentChatEvent struct { Skills []SkillInfo `json:"skills,omitempty"` // resolved per-user skills } -// Dispatcher routes agent chat requests to the executor. -// Two implementations: LocalDispatcher (direct goroutine) and NATSDispatcher (queue). -type Dispatcher interface { - // Dispatch sends a chat message to an agent and returns immediately. - // The response is delivered asynchronously via the configured event delivery mechanism. - Dispatch(userID, agentName, message string) (messageID string, err error) - - // Start initializes the dispatcher (e.g., subscribes to NATS queue). - Start(ctx context.Context) error -} - // ConfigProvider loads agent configs. Implemented by both file-based and DB-backed stores. type ConfigProvider interface { GetAgentConfig(userID, name string) (*AgentConfig, error) @@ -222,102 +211,64 @@ func (d *LocalDispatcher) buildLocalCallbacks(writer SSEWriter, messageID string // --- NATS Dispatcher (distributed) --- -// NATSDispatcher dispatches agent chats via NATS queue group. +// NATSDispatcher runs the agent chats a WorkConsumer delivers. type NATSDispatcher struct { - nats messaging.MessagingClient - eventBridge *EventBridge - configs ConfigProvider - apiURL string - apiKey string - subject string - queue string - sub messaging.Subscription // stored subscription for cleanup - sem chan struct{} // concurrency limiter; nil = unlimited - wg sync.WaitGroup + consumer messaging.WorkConsumer + eventBridge *EventBridge + configs ConfigProvider + apiURL string + apiKey string + maxConcurrent int + sub messaging.Subscription // stored subscription for cleanup } -// NewNATSDispatcher creates a dispatcher that uses NATS for distribution. -// maxConcurrent limits the number of concurrent agent jobs; 0 means unlimited. -func NewNATSDispatcher(nats messaging.MessagingClient, bridge *EventBridge, configs ConfigProvider, apiURL, apiKey, subject, queue string, maxConcurrent int) *NATSDispatcher { - d := &NATSDispatcher{ - nats: nats, - eventBridge: bridge, - configs: configs, - apiURL: apiURL, - apiKey: apiKey, - subject: subject, - queue: queue, +// NewNATSDispatcher creates a dispatcher that runs the agent runs consumer +// delivers. maxConcurrent limits the number of concurrent agent jobs; 0 means +// unlimited. +func NewNATSDispatcher(consumer messaging.WorkConsumer, bridge *EventBridge, configs ConfigProvider, apiURL, apiKey string, maxConcurrent int) *NATSDispatcher { + return &NATSDispatcher{ + consumer: consumer, + eventBridge: bridge, + configs: configs, + apiURL: apiURL, + apiKey: apiKey, + maxConcurrent: maxConcurrent, } - if maxConcurrent > 0 { - d.sem = make(chan struct{}, maxConcurrent) - } - return d } func (d *NATSDispatcher) Start(ctx context.Context) error { - sub, err := d.nats.QueueSubscribe(d.subject, d.queue, func(data []byte) { - var evt AgentChatEvent - if err := json.Unmarshal(data, &evt); err != nil { - xlog.Error("Failed to unmarshal agent chat event", "error", err) - return - } - if d.sem != nil { - select { - case d.sem <- struct{}{}: - case <-ctx.Done(): - return - } - } - d.wg.Add(1) - concurrency.SafeGo(func() { - defer d.wg.Done() - if d.sem != nil { - defer func() { <-d.sem }() - } - d.handleJob(ctx, evt) - }) - }) + sub, err := d.consumer.Consume(ctx, messaging.WorkAgentRun, d.maxConcurrent, d.runDelivery) if err != nil { - return fmt.Errorf("subscribing to %s: %w", d.subject, err) + return err } d.sub = sub - xlog.Info("NATS agent dispatcher started", "subject", d.subject, "queue", d.queue) + xlog.Info("NATS agent dispatcher started") return nil } -// Stop unsubscribes from the NATS queue, stopping message delivery. +// runDelivery ignores events: on NATS it is the same bus the process-wide +// event bridge already publishes on. An undecodable event returns nil because +// a carrier that redelivers on error would hand it back forever. +func (d *NATSDispatcher) runDelivery(ctx context.Context, payload []byte, _ messaging.Publisher) error { + var evt AgentChatEvent + if err := json.Unmarshal(payload, &evt); err != nil { + xlog.Error("Failed to unmarshal agent chat event", "error", err) + return nil + } + d.handleJob(ctx, evt) + return nil +} + +// Stop stops delivery and waits for the agent runs already in flight. func (d *NATSDispatcher) Stop() error { if d.sub != nil { err := d.sub.Unsubscribe() d.sub = nil - d.wg.Wait() return err } return nil } -func (d *NATSDispatcher) Dispatch(userID, agentName, message string) (string, error) { - messageID := uuid.New().String() - - // Send user message to SSE immediately - if d.eventBridge != nil { - d.eventBridge.PublishMessage(agentName, userID, RoleUser, message, messageID+"-user") - d.eventBridge.PublishStatus(agentName, userID, "processing") - } - - evt := AgentChatEvent{ - AgentName: agentName, - UserID: userID, - Message: message, - MessageID: messageID, - Role: RoleUser, - } - if err := d.nats.Publish(d.subject, evt); err != nil { - return "", fmt.Errorf("failed to dispatch agent chat: %w", err) - } - return messageID, nil -} - func (d *NATSDispatcher) handleJob(ctx context.Context, evt AgentChatEvent) { xlog.Info("Processing agent chat job", "agent", evt.AgentName, "user", evt.UserID) diff --git a/core/services/agents/dispatcher_test.go b/core/services/agents/dispatcher_test.go new file mode 100644 index 000000000..4e4a648e0 --- /dev/null +++ b/core/services/agents/dispatcher_test.go @@ -0,0 +1,159 @@ +package agents + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "log/slog" + "sync" + + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/xlog" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// recordingConfigProvider counts lookups: the dispatcher asks for a config +// only once it has decoded an event and is about to run the agent. +type recordingConfigProvider struct { + mu sync.Mutex + calls []string +} + +func (p *recordingConfigProvider) GetAgentConfig(userID, name string) (*AgentConfig, error) { + p.mu.Lock() + defer p.mu.Unlock() + p.calls = append(p.calls, userID+"/"+name) + return nil, errors.New("no such agent") +} + +func (p *recordingConfigProvider) Calls() []string { + p.mu.Lock() + defer p.mu.Unlock() + return append([]string(nil), p.calls...) +} + +// lockedBuffer lets the handler goroutine log while the spec reads. +type lockedBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *lockedBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *lockedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +var _ = Describe("NATSDispatcher consuming agent runs", func() { + var ( + bus *testutil.FakeBus + configs *recordingConfigProvider + logs *lockedBuffer + d *NATSDispatcher + ) + + BeforeEach(func() { + bus = testutil.NewFakeBus() + configs = &recordingConfigProvider{} + logs = &lockedBuffer{} + xlog.SetLogger(xlog.NewLoggerWithHandler(slog.NewTextHandler(logs, &slog.HandlerOptions{Level: slog.LevelError}), xlog.LogLevelError)) + DeferCleanup(func() { + xlog.SetLogger(xlog.NewLogger(xlog.LogLevel("info"), "text")) + }) + + bridge := NewEventBridge(bus, nil, "test-worker") + d = NewNATSDispatcher(messaging.NewNATSWorkConsumer(bus), bridge, configs, "http://127.0.0.1:1", "", 0) + Expect(d.Start(GinkgoT().Context())).To(Succeed()) + }) + + It("logs and drops an undecodable event without running an agent", func() { + var statuses []string + var mu sync.Mutex + _, err := bus.Subscribe("agent.*.events.*", func(data []byte) { + mu.Lock() + statuses = append(statuses, string(data)) + mu.Unlock() + }) + Expect(err).ToNot(HaveOccurred()) + + // A JSON string is valid on the wire but cannot decode into an event. + Expect(bus.Publish(messaging.SubjectAgentExecute, "not an event")).To(Succeed()) + // Stop waits for the in-flight handler, so everything below is final. + Expect(d.Stop()).To(Succeed()) + + Expect(logs.String()).To(ContainSubstring("Failed to unmarshal agent chat event")) + Expect(configs.Calls()).To(BeEmpty()) + mu.Lock() + defer mu.Unlock() + Expect(statuses).To(BeEmpty()) + }) + + It("runs a decodable event on the process-wide event bridge", func() { + var events []AgentEvent + var mu sync.Mutex + _, err := bus.Subscribe(messaging.SubjectAgentEvents("a1", "u1"), func(data []byte) { + var evt AgentEvent + Expect(json.Unmarshal(data, &evt)).To(Succeed()) + mu.Lock() + events = append(events, evt) + mu.Unlock() + }) + Expect(err).ToNot(HaveOccurred()) + + Expect(bus.Publish(messaging.SubjectAgentExecute, AgentChatEvent{AgentName: "a1", UserID: "u1", Message: "hi"})).To(Succeed()) + Expect(d.Stop()).To(Succeed()) + + Expect(configs.Calls()).To(Equal([]string{"u1/a1"})) + mu.Lock() + defer mu.Unlock() + Expect(events).To(HaveLen(1)) + Expect(events[0].EventType).To(Equal("json_message_status")) + Expect(events[0].Metadata).To(ContainSubstring("error: agent config not found")) + }) +}) + +// recordingWorkConsumer records what each Consume call asked for, so a spec +// can pin the limit the production caller chooses rather than what the +// carrier does with it. +type recordingWorkConsumer struct { + kinds []messaging.WorkKind + max []int +} + +func (c *recordingWorkConsumer) Consume(_ context.Context, kind messaging.WorkKind, maxInFlight int, _ messaging.WorkHandler) (messaging.Subscription, error) { + c.kinds = append(c.kinds, kind) + c.max = append(c.max, maxInFlight) + return noopSubscription{}, nil +} + +type noopSubscription struct{} + +func (noopSubscription) Unsubscribe() error { return nil } + +var _ = Describe("NATSDispatcher.Start", func() { + // The CLI agent worker passes 0, so agent runs are unbounded per worker; + // a dispatcher that dropped or replaced its limit would change that. + DescribeTable("asks for agent runs with its own concurrency limit", + func(maxConcurrent int) { + consumer := &recordingWorkConsumer{} + d := NewNATSDispatcher(consumer, nil, nil, "", "", maxConcurrent) + Expect(d.Start(GinkgoT().Context())).To(Succeed()) + + Expect(consumer.kinds).To(Equal([]messaging.WorkKind{messaging.WorkAgentRun})) + Expect(consumer.max).To(Equal([]int{maxConcurrent})) + }, + Entry("unbounded", 0), + Entry("serial", 1), + Entry("bounded", 4), + ) +}) diff --git a/core/services/agents/scheduler.go b/core/services/agents/scheduler.go index e159d8732..3bbc0b070 100644 --- a/core/services/agents/scheduler.go +++ b/core/services/agents/scheduler.go @@ -18,15 +18,14 @@ type SchedulerStore interface { } // AgentScheduler periodically checks for agents with standalone_job=true -// and publishes background run events to the NATS agent execution queue. +// and enqueues background run events as agent-run work. // Uses a PostgreSQL advisory lock so only one instance fires the cron. // Same pattern as notetaker's runAgentScheduler and LocalAI's cronLeaderLoop. type AgentScheduler struct { db *gorm.DB - nats messaging.Publisher + queue messaging.WorkQueue store SchedulerStore skillProvider SkillContentProvider // optional: loads full skill info for enriching events - subject string // NATS subject for agent execution pollInterval time.Duration // how often to check for due agents } @@ -41,12 +40,11 @@ func WithSchedulerSkillProvider(provider SkillContentProvider) AgentSchedulerOpt } // NewAgentScheduler creates a new background agent scheduler. -func NewAgentScheduler(db *gorm.DB, nats messaging.Publisher, store SchedulerStore, subject string, opts ...AgentSchedulerOpt) *AgentScheduler { +func NewAgentScheduler(db *gorm.DB, queue messaging.WorkQueue, store SchedulerStore, opts ...AgentSchedulerOpt) *AgentScheduler { s := &AgentScheduler{ db: db, - nats: nats, + queue: queue, store: store, - subject: subject, pollInterval: 15 * time.Second, } for _, opt := range opts { @@ -57,14 +55,14 @@ func NewAgentScheduler(db *gorm.DB, nats messaging.Publisher, store SchedulerSto // Start begins the scheduler loop. Blocks until ctx is cancelled. func (s *AgentScheduler) Start(ctx context.Context) { - xlog.Info("Agent scheduler started", "pollInterval", s.pollInterval, "subject", s.subject) - advisorylock.RunLeaderLoop(ctx, s.db, advisorylock.KeyAgentScheduler, s.pollInterval, s.runDueAgents) + xlog.Info("Agent scheduler started", "pollInterval", s.pollInterval) + advisorylock.RunLeaderLoop(ctx, s.db, advisorylock.KeyAgentScheduler, s.pollInterval, func() { s.runDueAgents(ctx) }) xlog.Info("Agent scheduler stopped") } // runDueAgents finds all agents with standalone_job=true that are due for a run -// and publishes background execution events to the NATS queue. -func (s *AgentScheduler) runDueAgents() { +// and enqueues background execution events. +func (s *AgentScheduler) runDueAgents(ctx context.Context) { configs, err := s.store.ListConfigs("") // all users if err != nil { xlog.Error("Agent scheduler: failed to list configs", "error", err) @@ -103,7 +101,6 @@ func (s *AgentScheduler) runDueAgents() { } } - // Publish background run event evt := AgentChatEvent{ AgentName: rec.Name, UserID: rec.UserID, @@ -112,8 +109,8 @@ func (s *AgentScheduler) runDueAgents() { Config: &cfg, Skills: skills, } - if err := s.nats.Publish(s.subject, evt); err != nil { - xlog.Error("Agent scheduler: failed to publish event", "agent", rec.Name, "error", err) + if err := s.queue.Enqueue(ctx, messaging.WorkAgentRun, evt); err != nil { + xlog.Error("Agent scheduler: failed to enqueue event", "agent", rec.Name, "error", err) continue } diff --git a/core/services/agents/scheduler_test.go b/core/services/agents/scheduler_test.go index 03e81690f..928323de2 100644 --- a/core/services/agents/scheduler_test.go +++ b/core/services/agents/scheduler_test.go @@ -1,27 +1,31 @@ package agents import ( + "context" "encoding/json" "time" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" + "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/testutil" ) -// mockPublisher records all Publish calls for assertions. -type mockPublisher struct { - calls []publishCall +// enqueueCall records one Enqueue with the typed payload, so a spec can tell +// an AgentChatEvent from pre-encoded bytes. +type enqueueCall struct { + kind messaging.WorkKind + payload any } -type publishCall struct { - subject string - data any +// fakeWorkQueue implements messaging.WorkQueue and records every Enqueue. +type fakeWorkQueue struct { + calls []enqueueCall } -func (m *mockPublisher) Publish(subject string, data any) error { - m.calls = append(m.calls, publishCall{subject: subject, data: data}) +func (f *fakeWorkQueue) Enqueue(_ context.Context, kind messaging.WorkKind, payload any) error { + f.calls = append(f.calls, enqueueCall{kind: kind, payload: payload}) return nil } @@ -120,19 +124,19 @@ var _ = Describe("AgentScheduler", func() { // ----------------------------------------------------------------------- Describe("runDueAgents", func() { var ( - pub *mockPublisher + queue *fakeWorkQueue mStore *mockSchedulerStore sched *AgentScheduler ) BeforeEach(func() { db := testutil.SetupTestDB() - pub = &mockPublisher{} + queue = &fakeWorkQueue{} mStore = &mockSchedulerStore{} - sched = NewAgentScheduler(db, pub, mStore, "agent.execute") + sched = NewAgentScheduler(db, queue, mStore) }) - It("publishes event for a due standalone agent", func() { + It("enqueues an agent run for a due standalone agent", func() { past := time.Now().Add(-15 * time.Minute) cfg := AgentConfig{ StandaloneJob: true, @@ -153,12 +157,12 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(HaveLen(1)) - Expect(pub.calls[0].subject).To(Equal("agent.execute")) + Expect(queue.calls).To(HaveLen(1)) + Expect(queue.calls[0].kind).To(Equal(messaging.WorkAgentRun)) - evt, ok := pub.calls[0].data.(AgentChatEvent) + evt, ok := queue.calls[0].payload.(AgentChatEvent) Expect(ok).To(BeTrue()) Expect(evt.AgentName).To(Equal("background-agent")) Expect(evt.UserID).To(Equal("user-1")) @@ -186,9 +190,9 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(BeEmpty()) + Expect(queue.calls).To(BeEmpty()) }) It("skips non-standalone agents", func() { @@ -210,9 +214,9 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(BeEmpty()) + Expect(queue.calls).To(BeEmpty()) }) It("skips paused agents", func() { @@ -234,9 +238,9 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(BeEmpty()) + Expect(queue.calls).To(BeEmpty()) }) It("skips agents with invalid config JSON", func() { @@ -253,12 +257,12 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(BeEmpty()) + Expect(queue.calls).To(BeEmpty()) }) - It("updates last run timestamp after publishing", func() { + It("updates last run timestamp after enqueueing", func() { cfg := AgentConfig{ StandaloneJob: true, PeriodicRuns: "10m", @@ -276,9 +280,9 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(HaveLen(1)) + Expect(queue.calls).To(HaveLen(1)) Expect(mStore.updated).To(HaveLen(1)) Expect(mStore.updated[0].userID).To(Equal("user-1")) Expect(mStore.updated[0].name).To(Equal("track-agent")) @@ -312,10 +316,10 @@ var _ = Describe("AgentScheduler", func() { } sched.skillProvider = provider - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(HaveLen(1)) - evt, ok := pub.calls[0].data.(AgentChatEvent) + Expect(queue.calls).To(HaveLen(1)) + evt, ok := queue.calls[0].payload.(AgentChatEvent) Expect(ok).To(BeTrue()) Expect(evt.Skills).To(HaveLen(2)) Expect(evt.Skills[0].Name).To(Equal("search")) @@ -349,12 +353,12 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(HaveLen(2)) + Expect(queue.calls).To(HaveLen(2)) names := []string{ - pub.calls[0].data.(AgentChatEvent).AgentName, - pub.calls[1].data.(AgentChatEvent).AgentName, + queue.calls[0].payload.(AgentChatEvent).AgentName, + queue.calls[1].payload.(AgentChatEvent).AgentName, } Expect(names).To(ConsistOf("agent-a", "agent-b")) }) @@ -379,9 +383,9 @@ var _ = Describe("AgentScheduler", func() { }, } - sched.runDueAgents() + sched.runDueAgents(context.Background()) - Expect(pub.calls).To(HaveLen(1)) + Expect(queue.calls).To(HaveLen(1)) }) }) }) diff --git a/core/services/failover/distsync/distsync.go b/core/services/failover/distsync/distsync.go index 825b495c6..f93ec93e1 100644 --- a/core/services/failover/distsync/distsync.go +++ b/core/services/failover/distsync/distsync.go @@ -41,7 +41,7 @@ type Sync struct { // New builds and starts the three maps, then attaches the result to m via // SetStateSync so any already-durable pins hydrate onto m immediately. -func New(ctx context.Context, nats messaging.MessagingClient, pins syncstate.Store[string, PinRecord], m *failover.Manager) (*Sync, error) { +func New(ctx context.Context, nats messaging.Broadcaster, pins syncstate.Store[string, PinRecord], m *failover.Manager) (*Sync, error) { s := &Sync{} // pins is already typed as the Store interface (the brief fixes this diff --git a/core/services/finetune/service.go b/core/services/finetune/service.go index 3e2431df2..e766ef43d 100644 --- a/core/services/finetune/service.go +++ b/core/services/finetune/service.go @@ -52,7 +52,7 @@ func NewFineTuneService( appConfig *config.ApplicationConfig, modelLoader *model.ModelLoader, configLoader *config.ModelConfigLoader, - nats messaging.MessagingClient, + nats messaging.Broadcaster, store *distributed.FineTuneStore, ) *FineTuneService { s := &FineTuneService{ diff --git a/core/services/galleryop/operation.go b/core/services/galleryop/operation.go index 4322b626d..1bfee6420 100644 --- a/core/services/galleryop/operation.go +++ b/core/services/galleryop/operation.go @@ -237,7 +237,7 @@ type OpCache struct { // Distributed sync (nil when standalone). mu sync.RWMutex - nats messaging.MessagingClient + nats messaging.Broadcaster store *distributed.GalleryStore subs []messaging.Subscription } @@ -255,7 +255,7 @@ func NewOpCache(galleryService *GalleryService) *OpCache { // SetMessagingClient enables cross-replica OpCache sync. Once set, Set/ // SetBackend/DeleteUUID publish OpCacheEvent messages that peer OpCaches // merge into their local maps. Call Start after this to subscribe. -func (m *OpCache) SetMessagingClient(nc messaging.MessagingClient) { +func (m *OpCache) SetMessagingClient(nc messaging.Broadcaster) { m.mu.Lock() defer m.mu.Unlock() m.nats = nc diff --git a/core/services/galleryop/service.go b/core/services/galleryop/service.go index 6d6da5d1e..68dbca782 100644 --- a/core/services/galleryop/service.go +++ b/core/services/galleryop/service.go @@ -31,10 +31,10 @@ type GalleryService struct { cancellations map[string]cancellationActions // Distributed mode (nil when not in distributed mode). - // natsClient is the wider MessagingClient (Publisher + subscribe methods) + // natsClient is a messaging.Broadcaster (Publisher + Subscribe) // when wired by the distributed startup path; broadcastSubs holds the // progress + cancel subscriptions opened by SubscribeBroadcasts. - natsClient messaging.MessagingClient + natsClient messaging.Broadcaster galleryStore *distributed.GalleryStore broadcastSubs []messaging.Subscription @@ -118,10 +118,10 @@ func (g *GalleryService) ModelArtifactMaterializer() config.ArtifactMaterializer } // SetNATSClient sets the NATS client for distributed progress publishing. -// Accepting the wider MessagingClient (vs. plain Publisher) lets +// Accepting a Broadcaster (vs. plain Publisher) lets // SubscribeBroadcasts wire the wildcard subscriptions that keep peer // replicas' statuses + cancellations in sync. -func (g *GalleryService) SetNATSClient(nc messaging.MessagingClient) { +func (g *GalleryService) SetNATSClient(nc messaging.Broadcaster) { g.Lock() defer g.Unlock() g.natsClient = nc diff --git a/core/services/jobs/dispatcher.go b/core/services/jobs/dispatcher.go index a1da792ed..6f0a8f568 100644 --- a/core/services/jobs/dispatcher.go +++ b/core/services/jobs/dispatcher.go @@ -2,13 +2,11 @@ package jobs import ( "context" - "errors" "fmt" "time" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/services/advisorylock" - "github.com/mudler/LocalAI/pkg/concurrency" "github.com/mudler/LocalAI/core/services/dbutil" "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/xlog" @@ -55,53 +53,36 @@ type CancelEvent struct { JobID string `json:"job_id"` } -// WorkerFunc is the function signature for processing a job. -// It receives the job record, task record, and a context that will be cancelled -// if the job is cancelled via NATS. -type WorkerFunc func(ctx context.Context, job *JobRecord, task *TaskRecord) error - -// Dispatcher distributes jobs across instances via NATS queue groups -// and coordinates cron execution via PostgreSQL advisory locks. +// Dispatcher hands jobs to the work queue, persists the results and traces +// workers publish, and coordinates cron execution via PostgreSQL advisory locks. type Dispatcher struct { store *JobStore - nats messaging.MessagingClient + queue messaging.WorkQueue + nats messaging.Broadcaster db *gorm.DB instanceID string configLoader ModelConfigLoader // optional: to enrich job events with model config - // Worker function (set by the application) - workerFn WorkerFunc - - // Cancel registry (notetaker pattern) - cancelRegistry messaging.CancelRegistry - // NATS subscriptions - jobSub messaging.Subscription - cancelSub messaging.Subscription resultSub messaging.Subscription progressSub messaging.Subscription - // Concurrency limiter; nil = unlimited - sem chan struct{} - // Lifecycle ctx context.Context cancel context.CancelFunc } -// NewDispatcher creates a new distributed job Dispatcher. -// maxConcurrent limits the number of concurrent job goroutines; 0 means unlimited. -func NewDispatcher(store *JobStore, nc messaging.MessagingClient, db *gorm.DB, instanceID string, maxConcurrent int) *Dispatcher { - d := &Dispatcher{ +// NewDispatcher creates a new distributed job Dispatcher. Jobs leave through +// queue and workers consume them; nc carries cancel, progress and result +// fan-out. +func NewDispatcher(store *JobStore, queue messaging.WorkQueue, nc messaging.Broadcaster, db *gorm.DB, instanceID string) *Dispatcher { + return &Dispatcher{ store: store, + queue: queue, nats: nc, db: db, instanceID: instanceID, } - if maxConcurrent > 0 { - d.sem = make(chan struct{}, maxConcurrent) - } - return d } // ModelConfigLoader loads model configurations by name. @@ -109,29 +90,13 @@ type ModelConfigLoader interface { GetModelConfig(name string) (config.ModelConfig, bool) } -// NewWorkerDispatcher creates a dispatcher that also consumes and processes jobs. -// Use this instead of NewDispatcher + SetWorkerFunc + SetModelConfigLoader when both -// the worker function and config loader are available at construction time. -func NewWorkerDispatcher(store *JobStore, nc messaging.MessagingClient, db *gorm.DB, instanceID string, maxConcurrent int, workerFn WorkerFunc, configLoader ModelConfigLoader) *Dispatcher { - d := NewDispatcher(store, nc, db, instanceID, maxConcurrent) - d.workerFn = workerFn - d.configLoader = configLoader - return d -} - -// SetWorkerFunc sets the function that processes jobs. -// Deprecated: prefer NewWorkerDispatcher when the worker function is available at construction time. -func (d *Dispatcher) SetWorkerFunc(fn WorkerFunc) { - d.workerFn = fn -} - // SetModelConfigLoader sets the model config loader for enriching job events. -// Deprecated: prefer NewWorkerDispatcher when the config loader is available at construction time. func (d *Dispatcher) SetModelConfigLoader(cl ModelConfigLoader) { d.configLoader = cl } -// Start begins listening for jobs via NATS and starts the cron leader loop. +// Start subscribes to the results and traces workers publish and starts the +// cron leader loop. It consumes no jobs: workers do. func (d *Dispatcher) Start(ctx context.Context) error { d.ctx, d.cancel = context.WithCancel(ctx) success := false @@ -141,35 +106,7 @@ func (d *Dispatcher) Start(ctx context.Context) error { } }() - // Subscribe to job queue only if a worker function is configured. - // In distributed mode, the frontend dispatcher publishes jobs but does not consume them — - // agent workers pick them up from the same NATS queue. var err error - if d.workerFn != nil { - d.jobSub, err = messaging.QueueSubscribeJSON(d.nats, messaging.SubjectJobsNew, messaging.QueueWorkers, func(evt JobEvent) { - concurrency.SafeGo(func() { - if d.sem != nil { - d.sem <- struct{}{} - defer func() { <-d.sem }() - } - d.processJob(evt) - }) - }) - if err != nil { - return fmt.Errorf("subscribing to job queue: %w", err) - } - } - - // Subscribe to cancel events (broadcast to all — each instance checks its registry) - d.cancelSub, err = messaging.SubscribeJSON(d.nats, messaging.SubjectJobCancelWildcard, func(evt CancelEvent) { - if d.cancelRegistry.Cancel(evt.JobID) { - xlog.Info("Cancelled job via NATS", "jobID", evt.JobID) - } - }) - if err != nil { - return fmt.Errorf("subscribing to cancel events: %w", err) - } - // Subscribe to job result events from workers (persist to DB) if d.store != nil { d.resultSub, err = messaging.SubscribeJSON(d.nats, messaging.SubjectJobResultWildcard, func(evt JobResultEvent) { @@ -203,14 +140,6 @@ func (d *Dispatcher) Start(ctx context.Context) error { // unsubscribeAll nil-checks, unsubscribes, and nils out each NATS subscription. // Safe to call multiple times. func (d *Dispatcher) unsubscribeAll() { - if d.jobSub != nil { - d.jobSub.Unsubscribe() - d.jobSub = nil - } - if d.cancelSub != nil { - d.cancelSub.Unsubscribe() - d.cancelSub = nil - } if d.resultSub != nil { d.resultSub.Unsubscribe() d.resultSub = nil @@ -221,7 +150,7 @@ func (d *Dispatcher) unsubscribeAll() { } } -// Stop cleans up subscriptions and cancels running jobs. +// Stop cleans up subscriptions and stops the cron leader loop. func (d *Dispatcher) Stop() { if d.cancel != nil { d.cancel() @@ -229,7 +158,7 @@ func (d *Dispatcher) Stop() { d.unsubscribeAll() } -// Enqueue publishes a job to the NATS queue for distributed processing. +// Enqueue hands a job to the work queue for distributed processing. // The event is enriched with the full Job and Task records so that the // worker does not need direct database access. func (d *Dispatcher) Enqueue(jobID, taskID, userID string) error { @@ -255,12 +184,14 @@ func (d *Dispatcher) Enqueue(jobID, taskID, userID string) error { } } - subject := messaging.SubjectJobsNew + kind := messaging.WorkTask if evt.ModelConfig != nil && evt.ModelConfig.MCP.HasMCPServers() { - subject = messaging.SubjectMCPCIJobsNew + kind = messaging.WorkMCPCI } - return d.nats.Publish(subject, evt) + // Enqueue takes no ctx from its callers (an HTTP handler and the cron + // loop), and the NATS carrier ignores it anyway. + return d.queue.Enqueue(context.Background(), kind, evt) } // Cancel publishes a cancel event to NATS (broadcast to all instances). @@ -284,116 +215,6 @@ func (d *Dispatcher) SubscribeProgress(jobID string, handler func(ProgressEvent) return messaging.SubscribeJSON(d.nats, messaging.SubjectJobProgress(jobID), handler) } -// processJob is called by the NATS queue subscriber to execute a job. -// It prefers Job+Task from the enriched NATS payload (no DB needed). -// Results are published back via NATS for the frontend to persist. -func (d *Dispatcher) processJob(evt JobEvent) { - if d.workerFn == nil { - xlog.Error("No worker function set for job dispatcher") - d.publishResult(evt.JobID, "failed", "", "no worker function configured") - return - } - - // Prefer enriched payload; fall back to DB for backward compat - job := evt.Job - if job == nil && d.store != nil { - var err error - job, err = d.store.GetJob(evt.JobID) - if err != nil { - xlog.Error("Failed to load job", "jobID", evt.JobID, "error", err) - return - } - } - if job == nil { - xlog.Error("No job data available", "jobID", evt.JobID) - return - } - - task := evt.Task - if task == nil && d.store != nil { - var err error - task, err = d.store.GetTask(job.TaskID) - if err != nil { - xlog.Error("Failed to load task for job", "jobID", evt.JobID, "taskID", job.TaskID, "error", err) - d.publishResult(evt.JobID, "failed", "", "task not found") - return - } - } - if task == nil { - xlog.Error("No task data available", "jobID", evt.JobID) - d.publishResult(evt.JobID, "failed", "", "task not found") - return - } - - // Pre-register so cancels arriving before context creation are captured - cancelled := make(chan struct{}, 1) - d.cancelRegistry.Register(evt.JobID, func() { - select { - case cancelled <- struct{}{}: - default: - } - }) - - ctx, cancelFn := context.WithCancel(d.ctx) - d.cancelRegistry.Register(evt.JobID, cancelFn) // overwrite with real cancel - - // Check if cancel arrived during the registration window - select { - case <-cancelled: - cancelFn() - default: - } - - // Check if job was cancelled in the DB before we picked it up - if d.store != nil { - if dbJob, err := d.store.GetJob(evt.JobID); err == nil && dbJob.Status == "cancelled" { - cancelFn() - } - } - - defer func() { - d.cancelRegistry.Deregister(evt.JobID) - cancelFn() - }() - - // Check if already cancelled before starting - select { - case <-ctx.Done(): - d.publishResult(evt.JobID, "cancelled", "", "") - return - default: - } - - // Mark as running - job.FrontendID = d.instanceID - d.PublishProgress(evt.JobID, "running", "Job started") - - // Execute - err := d.workerFn(ctx, job, task) - - if errors.Is(ctx.Err(), context.Canceled) { - d.publishResult(evt.JobID, "cancelled", "", "") - d.PublishProgress(evt.JobID, "cancelled", "Job cancelled") - return - } - - if err != nil { - d.publishResult(evt.JobID, "failed", "", err.Error()) - d.PublishProgress(evt.JobID, "failed", err.Error()) - return - } - - // Publish completion — result is set on the job by workerFn - d.publishResult(evt.JobID, "completed", job.Result, "") - d.PublishProgress(evt.JobID, "completed", "Job completed") -} - -// publishResult publishes the terminal job result via NATS. -// The frontend subscribes to these events and persists to DB. -func (d *Dispatcher) publishResult(jobID, status, result, errMsg string) { - PublishJobResult(d.nats, jobID, status, result, errMsg) -} - // PublishTrace publishes a trace event for a running job via NATS. // The frontend subscribes and persists traces to DB. func (d *Dispatcher) PublishTrace(jobID, traceType, traceContent string) error { diff --git a/core/services/jobs/dispatcher_test.go b/core/services/jobs/dispatcher_test.go index 0af251fdd..dbf27c846 100644 --- a/core/services/jobs/dispatcher_test.go +++ b/core/services/jobs/dispatcher_test.go @@ -1,6 +1,7 @@ package jobs import ( + "context" "encoding/json" "time" @@ -12,50 +13,23 @@ import ( "github.com/mudler/LocalAI/core/services/testutil" ) -// publishCall records a single Publish invocation. -type publishCall struct { - subject string - data any +// enqueueCall records a single Enqueue invocation with the typed payload, so +// a spec can tell a JobEvent from pre-encoded bytes. +type enqueueCall struct { + kind messaging.WorkKind + payload any } -// fakeMessagingClient implements messaging.MessagingClient and records published messages. -type fakeMessagingClient struct { - calls []publishCall +// fakeWorkQueue implements messaging.WorkQueue and records every Enqueue. +type fakeWorkQueue struct { + calls []enqueueCall } -func (f *fakeMessagingClient) Publish(subject string, data any) error { - f.calls = append(f.calls, publishCall{subject: subject, data: data}) +func (f *fakeWorkQueue) Enqueue(_ context.Context, kind messaging.WorkKind, payload any) error { + f.calls = append(f.calls, enqueueCall{kind: kind, payload: payload}) return nil } -func (f *fakeMessagingClient) Subscribe(string, func([]byte)) (messaging.Subscription, error) { - return &fakeSub{}, nil -} - -func (f *fakeMessagingClient) QueueSubscribe(string, string, func([]byte)) (messaging.Subscription, error) { - return &fakeSub{}, nil -} - -func (f *fakeMessagingClient) QueueSubscribeReply(string, string, func([]byte, func([]byte))) (messaging.Subscription, error) { - return &fakeSub{}, nil -} - -func (f *fakeMessagingClient) SubscribeReply(string, func([]byte, func([]byte))) (messaging.Subscription, error) { - return &fakeSub{}, nil -} - -func (f *fakeMessagingClient) Request(string, []byte, time.Duration) ([]byte, error) { - return nil, nil -} - -func (f *fakeMessagingClient) IsConnected() bool { return true } -func (f *fakeMessagingClient) Close() {} - -// fakeSub implements messaging.Subscription. -type fakeSub struct{} - -func (s *fakeSub) Unsubscribe() error { return nil } - // mockConfigLoader implements ModelConfigLoader for testing Enqueue routing. type mockConfigLoader struct { configs map[string]config.ModelConfig @@ -83,7 +57,7 @@ var _ = Describe("Dispatcher", func() { store, err = NewJobStore(db) Expect(err).ToNot(HaveOccurred()) - disp = NewDispatcher(store, nil, db, "test-instance", 0) + disp = NewDispatcher(store, nil, nil, db, "test-instance") }) It("returns true when no previous job exists", func() { @@ -209,12 +183,13 @@ var _ = Describe("Dispatcher", func() { }) // ----------------------------------------------------------------------- - // Enqueue — test NATS subject routing via real Dispatcher.Enqueue() + // Enqueue: the work kind chosen by the real Dispatcher.Enqueue() // ----------------------------------------------------------------------- - Describe("Enqueue subject routing", func() { + Describe("Enqueue work kind routing", func() { var ( store *JobStore - fake *fakeMessagingClient + queue *fakeWorkQueue + bus *testutil.FakeBus disp *Dispatcher ) @@ -223,11 +198,12 @@ var _ = Describe("Dispatcher", func() { var err error store, err = NewJobStore(db) Expect(err).ToNot(HaveOccurred()) - fake = &fakeMessagingClient{} - disp = NewDispatcher(store, fake, db, "test-instance", 0) + queue = &fakeWorkQueue{} + bus = testutil.NewFakeBus() + disp = NewDispatcher(store, queue, bus, db, "test-instance") }) - It("routes MCP jobs to SubjectMCPCIJobsNew", func() { + It("enqueues MCP jobs as WorkMCPCI", func() { task := &TaskRecord{ UserID: "user-1", Name: "mcp-task", @@ -256,11 +232,15 @@ var _ = Describe("Dispatcher", func() { Expect(disp.Enqueue(job.ID, task.ID, "user-1")).To(Succeed()) - Expect(fake.calls).To(HaveLen(1)) - Expect(fake.calls[0].subject).To(Equal(messaging.SubjectMCPCIJobsNew)) + Expect(queue.calls).To(HaveLen(1)) + Expect(queue.calls[0].kind).To(Equal(messaging.WorkMCPCI)) + evt, ok := queue.calls[0].payload.(JobEvent) + Expect(ok).To(BeTrue(), "the queue must be handed the typed JobEvent") + Expect(evt.JobID).To(Equal(job.ID)) + Expect(evt.TaskID).To(Equal(task.ID)) }) - It("routes non-MCP jobs to SubjectJobsNew", func() { + It("enqueues non-MCP jobs as WorkTask", func() { task := &TaskRecord{ UserID: "user-1", Name: "plain-task", @@ -285,11 +265,14 @@ var _ = Describe("Dispatcher", func() { Expect(disp.Enqueue(job.ID, task.ID, "user-1")).To(Succeed()) - Expect(fake.calls).To(HaveLen(1)) - Expect(fake.calls[0].subject).To(Equal(messaging.SubjectJobsNew)) + Expect(queue.calls).To(HaveLen(1)) + Expect(queue.calls[0].kind).To(Equal(messaging.WorkTask)) + evt, ok := queue.calls[0].payload.(JobEvent) + Expect(ok).To(BeTrue(), "the queue must be handed the typed JobEvent") + Expect(evt.JobID).To(Equal(job.ID)) }) - It("routes to SubjectJobsNew when model config is not found", func() { + It("enqueues as WorkTask when model config is not found", func() { task := &TaskRecord{ UserID: "user-1", Name: "unknown-model-task", @@ -312,18 +295,31 @@ var _ = Describe("Dispatcher", func() { Expect(disp.Enqueue(job.ID, task.ID, "user-1")).To(Succeed()) - Expect(fake.calls).To(HaveLen(1)) - Expect(fake.calls[0].subject).To(Equal(messaging.SubjectJobsNew)) + Expect(queue.calls).To(HaveLen(1)) + Expect(queue.calls[0].kind).To(Equal(messaging.WorkTask)) + }) + + It("keeps queued work off the fan-out bus", func() { + task := &TaskRecord{UserID: "user-1", Name: "bus-task", Model: "m", Enabled: true} + Expect(store.CreateTask(task)).To(Succeed()) + job := &JobRecord{TaskID: task.ID, UserID: "user-1", Status: "pending", TriggeredBy: "manual"} + Expect(store.CreateJob(job)).To(Succeed()) + + Expect(disp.Enqueue(job.ID, task.ID, "user-1")).To(Succeed()) + + Expect(queue.calls).To(HaveLen(1)) + Expect(bus.PublishCount(messaging.SubjectJobsNew)).To(BeZero()) + Expect(bus.PublishCount(messaging.SubjectMCPCIJobsNew)).To(BeZero()) }) }) // ----------------------------------------------------------------------- - // Enqueue event enrichment — verify the payload published by Enqueue() + // Enqueue event enrichment: verify the payload enqueued by Enqueue() // ----------------------------------------------------------------------- Describe("Enqueue event enrichment", func() { var ( store *JobStore - fake *fakeMessagingClient + queue *fakeWorkQueue disp *Dispatcher ) @@ -332,8 +328,8 @@ var _ = Describe("Dispatcher", func() { var err error store, err = NewJobStore(db) Expect(err).ToNot(HaveOccurred()) - fake = &fakeMessagingClient{} - disp = NewDispatcher(store, fake, db, "test-instance", 0) + queue = &fakeWorkQueue{} + disp = NewDispatcher(store, queue, nil, db, "test-instance") }) It("includes full job and task records in the event", func() { @@ -368,9 +364,9 @@ var _ = Describe("Dispatcher", func() { Expect(disp.Enqueue(job.ID, task.ID, "user-1")).To(Succeed()) - Expect(fake.calls).To(HaveLen(1)) - evt, ok := fake.calls[0].data.(JobEvent) - Expect(ok).To(BeTrue(), "published data should be a JobEvent") + Expect(queue.calls).To(HaveLen(1)) + evt, ok := queue.calls[0].payload.(JobEvent) + Expect(ok).To(BeTrue(), "enqueued payload should be a JobEvent") Expect(evt.Job).ToNot(BeNil()) Expect(evt.Job.ID).To(Equal(job.ID)) Expect(evt.Task).ToNot(BeNil()) @@ -399,8 +395,8 @@ var _ = Describe("Dispatcher", func() { // No config loader — Enqueue still works, just no model config enrichment. Expect(disp.Enqueue(job.ID, task.ID, "user-1")).To(Succeed()) - Expect(fake.calls).To(HaveLen(1)) - evt, ok := fake.calls[0].data.(JobEvent) + Expect(queue.calls).To(HaveLen(1)) + evt, ok := queue.calls[0].payload.(JobEvent) Expect(ok).To(BeTrue()) data, err := json.Marshal(evt) @@ -412,4 +408,67 @@ var _ = Describe("Dispatcher", func() { Expect(decoded.TaskID).To(Equal(task.ID)) }) }) + + // ----------------------------------------------------------------------- + // The frontend dispatcher only fans out: jobs leave through the WorkQueue + // and nothing here consumes them, so a Broadcaster is all it may ask for. + // ----------------------------------------------------------------------- + Describe("on a fan-out-only bus", func() { + var ( + store *JobStore + queue *fakeWorkQueue + bus *testutil.FakeBus + disp *Dispatcher + ) + + BeforeEach(func() { + db := testutil.SetupTestDB() + var err error + store, err = NewJobStore(db) + Expect(err).ToNot(HaveOccurred()) + queue = &fakeWorkQueue{} + bus = testutil.NewFakeBus() + // broadcastOnly hides the queue and request methods, so this + // compiles only while NewDispatcher asks for no more than it uses. + disp = NewDispatcher(store, queue, broadcastOnly{bus}, db, "test-instance") + ctx, cancel := context.WithCancel(context.Background()) + Expect(disp.Start(ctx)).To(Succeed()) + DeferCleanup(func() { + disp.Stop() + cancel() + }) + }) + + It("still enqueues jobs and joins no queue group", func() { + task := &TaskRecord{UserID: "user-1", Name: "fanout-task", Model: "m", Enabled: true} + Expect(store.CreateTask(task)).To(Succeed()) + job := &JobRecord{TaskID: task.ID, UserID: "user-1", Status: "pending", TriggeredBy: "manual"} + Expect(store.CreateJob(job)).To(Succeed()) + + Expect(disp.Enqueue(job.ID, task.ID, "user-1")).To(Succeed()) + + Expect(queue.calls).To(HaveLen(1)) + Expect(queue.calls[0].kind).To(Equal(messaging.WorkTask)) + Expect(bus.QueueGroups()).To(BeEmpty()) + }) + + It("persists the result a worker publishes", func() { + task := &TaskRecord{UserID: "user-1", Name: "result-task", Model: "m", Enabled: true} + Expect(store.CreateTask(task)).To(Succeed()) + job := &JobRecord{TaskID: task.ID, UserID: "user-1", Status: "running", TriggeredBy: "manual"} + Expect(store.CreateJob(job)).To(Succeed()) + + PublishJobResult(bus, job.ID, "completed", "the answer", "") + + stored, err := store.GetJob(job.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(stored.Status).To(Equal("completed")) + Expect(stored.Result).To(Equal("the answer")) + }) + }) }) + +// broadcastOnly narrows a FakeBus to the Broadcaster surface. +type broadcastOnly struct { + messaging.Broadcaster +} diff --git a/core/services/mcp/remote.go b/core/services/mcp/remote.go index 17cfc1f36..2b926e475 100644 --- a/core/services/mcp/remote.go +++ b/core/services/mcp/remote.go @@ -1,6 +1,8 @@ package mcp import ( + "context" + "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/pkg/functions" ) @@ -16,6 +18,15 @@ type MCPToolRequest struct { StdioServers config.MCPGenericConfig[config.MCPSTDIOServers] `json:"stdio_servers"` } +// ToolHandler serves one MCP tool request on an agent worker. It returns no +// error because every failure is an answer: the reply carries it in Error, so +// the requester never waits out its budget on silence. +type ToolHandler func(ctx context.Context, req MCPToolRequest) MCPToolResponse + +// DiscoveryHandler serves one MCP discovery request on an agent worker, with +// the same failure contract as ToolHandler. +type DiscoveryHandler func(ctx context.Context, req MCPDiscoveryRequest) MCPDiscoveryResponse + // MCPToolResponse is the NATS reply for an MCP tool execution. type MCPToolResponse struct { Result string `json:"result,omitempty"` diff --git a/core/services/messaging/backend_install_progress.go b/core/services/messaging/backend_install_progress.go index 268ef86b9..f0745d0b1 100644 --- a/core/services/messaging/backend_install_progress.go +++ b/core/services/messaging/backend_install_progress.go @@ -1,33 +1,5 @@ package messaging -// Phase values published on the BackendInstallProgressEvent.Phase field. -// Defined as exported constants so producer (worker install handler) and -// consumer (master bridge into OpStatus) share a single source of truth -// instead of two copies of the literal string. -const ( - PhaseResolving = "resolving" // worker is locating the gallery / image manifest - PhaseDownloading = "downloading" // worker is actively pulling layers - PhaseExtracting = "extracting" // worker is unpacking the downloaded archive - PhaseStarting = "starting" // worker is spawning the gRPC backend process -) - -// BackendInstallProgressEvent is the wire payload published by a worker to -// nodes..backend.install..progress while a long-running install -// is in flight. Transient: dropped events are acceptable, the master relies -// on BackendInstallReply for ground truth on success/failure. -// -// Phase holds one of the Phase* constants above. -type BackendInstallProgressEvent struct { - OpID string `json:"op_id"` - NodeID string `json:"node_id"` - Backend string `json:"backend"` - FileName string `json:"file_name,omitempty"` - Current string `json:"current,omitempty"` // human-readable size, e.g. "412 MB" - Total string `json:"total,omitempty"` // human-readable size, e.g. "2.1 GB" - Percentage float64 `json:"percentage"` - Phase string `json:"phase,omitempty"` -} - // SubjectNodeBackendInstallProgress returns the NATS subject for transient // progress events emitted by a worker during a single backend.install run. // Per-op so multiple concurrent installs on the same node never alias. diff --git a/core/services/messaging/backend_install_progress_test.go b/core/services/messaging/backend_install_progress_test.go index ec45f4619..5b5be4f3a 100644 --- a/core/services/messaging/backend_install_progress_test.go +++ b/core/services/messaging/backend_install_progress_test.go @@ -1,7 +1,6 @@ package messaging_test import ( - "encoding/json" "strings" . "github.com/onsi/ginkgo/v2" @@ -10,21 +9,6 @@ import ( "github.com/mudler/LocalAI/core/services/messaging" ) -var _ = Describe("Phase constants", func() { - // Pin the wire-format string values. A future refactor that renames - // a constant must NOT silently change the JSON value the master - // receives or break consumers that switch on Phase. - DescribeTable("phase constant", - func(actual, expected string) { - Expect(actual).To(Equal(expected)) - }, - Entry("resolving", messaging.PhaseResolving, "resolving"), - Entry("downloading", messaging.PhaseDownloading, "downloading"), - Entry("extracting", messaging.PhaseExtracting, "extracting"), - Entry("starting", messaging.PhaseStarting, "starting"), - ) -}) - var _ = Describe("BackendInstallProgress", func() { Context("SubjectNodeBackendInstallProgress", func() { It("composes the per-op progress subject", func() { @@ -42,25 +26,4 @@ var _ = Describe("BackendInstallProgress", func() { Expect(strings.Count(subj, ".")).To(Equal(5)) }) }) - - Context("BackendInstallProgressEvent", func() { - It("JSON round-trips with all known fields", func() { - ev := messaging.BackendInstallProgressEvent{ - OpID: "op-123", - NodeID: "node-abc", - Backend: "vllm", - FileName: "vllm-cpu.tar.zst", - Current: "412 MB", - Total: "2.1 GB", - Percentage: 19.6, - Phase: "downloading", - } - raw, err := json.Marshal(ev) - Expect(err).ToNot(HaveOccurred()) - - var got messaging.BackendInstallProgressEvent - Expect(json.Unmarshal(raw, &got)).To(Succeed()) - Expect(got).To(Equal(ev)) - }) - }) }) diff --git a/core/services/messaging/client.go b/core/services/messaging/client.go index e01c7d9ca..b8ab08059 100644 --- a/core/services/messaging/client.go +++ b/core/services/messaging/client.go @@ -147,6 +147,9 @@ func (c *Client) runReconnectCallbacks() { // Publish marshals data as JSON and publishes it to the given subject. func (c *Client) Publish(subject string, data any) error { + if err := ValidateSubject(subject); err != nil { + return err + } payload, err := json.Marshal(data) if err != nil { return fmt.Errorf("marshalling message for %s: %w", subject, err) @@ -184,6 +187,9 @@ func (c *Client) QueueSubscribe(subject, queue string, handler func([]byte)) (Su // lacks a subject gets a non-nil subscription that never receives a message, // turning a permission misconfiguration into a silent failure. func (c *Client) confirmSubscription(subject string, mk func(*nats.Conn) (*nats.Subscription, error)) (Subscription, error) { + if err := ValidateSubject(subject); err != nil { + return nil, err + } c.mu.RLock() conn := c.conn c.mu.RUnlock() @@ -222,6 +228,9 @@ func (c *Client) confirmSubscription(subject string, mk func(*nats.Conn) (*nats. // Request sends a request and waits for a reply (request-reply pattern). // Returns the raw reply data. func (c *Client) Request(subject string, data []byte, timeout time.Duration) ([]byte, error) { + if err := ValidateSubject(subject); err != nil { + return nil, err + } c.mu.RLock() defer c.mu.RUnlock() msg, err := c.conn.Request(subject, data, timeout) @@ -265,7 +274,7 @@ func (c *Client) QueueSubscribeReply(subject, queue string, handler func(data [] // SubscribeJSON creates a subscription that automatically unmarshals JSON messages. // Invalid JSON messages are logged and skipped. -func SubscribeJSON[T any](c MessagingClient, subject string, handler func(T)) (Subscription, error) { +func SubscribeJSON[T any](c Broadcaster, subject string, handler func(T)) (Subscription, error) { return c.Subscribe(subject, func(data []byte) { var evt T if err := json.Unmarshal(data, &evt); err != nil { @@ -276,19 +285,6 @@ func SubscribeJSON[T any](c MessagingClient, subject string, handler func(T)) (S }) } -// QueueSubscribeJSON creates a queue subscription that automatically unmarshals JSON messages. -// Invalid JSON messages are logged and skipped. -func QueueSubscribeJSON[T any](c MessagingClient, subject, queue string, handler func(T)) (Subscription, error) { - return c.QueueSubscribe(subject, queue, func(data []byte) { - var evt T - if err := json.Unmarshal(data, &evt); err != nil { - xlog.Warn("Failed to unmarshal NATS message", "subject", subject, "error", err) - return - } - handler(evt) - }) -} - // RequestJSON sends a JSON request-reply via NATS, marshaling the request and // unmarshaling the reply. This eliminates the repeated marshal/request/unmarshal // boilerplate across all NATS request-reply call sites. diff --git a/core/services/messaging/client_conformance_test.go b/core/services/messaging/client_conformance_test.go new file mode 100644 index 000000000..8ab456710 --- /dev/null +++ b/core/services/messaging/client_conformance_test.go @@ -0,0 +1,81 @@ +package messaging_test + +import ( + "context" + "fmt" + "os" + "runtime" + "sync" + + . "github.com/onsi/ginkgo/v2" + "github.com/testcontainers/testcontainers-go" + "github.com/testcontainers/testcontainers-go/wait" + + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/messaging/messagingtest" +) + +var ( + natsOnce sync.Once + natsURL string + natsCtr testcontainers.Container + natsErr error +) + +// sharedNATS starts one server for the whole suite. A container per spec would +// cost more than the suite itself. +func sharedNATS() (string, error) { + natsOnce.Do(func() { + ctx := context.Background() + natsCtr, natsErr = testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{ + ContainerRequest: testcontainers.ContainerRequest{ + Image: "nats:2.10-alpine", + ExposedPorts: []string{"4222/tcp"}, + WaitingFor: wait.ForListeningPort("4222/tcp"), + }, + Started: true, + }) + if natsErr != nil { + return + } + host, err := natsCtr.Host(ctx) + if err != nil { + natsErr = err + return + } + port, err := natsCtr.MappedPort(ctx, "4222/tcp") + if err != nil { + natsErr = err + return + } + natsURL = fmt.Sprintf("nats://%s:%s", host, port.Port()) + }) + return natsURL, natsErr +} + +var _ = AfterSuite(func() { + if natsCtr != nil { + _ = natsCtr.Terminate(context.Background()) + } +}) + +var _ = Describe("NATS client", func() { + messagingtest.RunBroadcasterConformance(func() (messaging.Broadcaster, func()) { + url, err := sharedNATS() + if err != nil { + // This is the only spec that runs the rules against a real carrier. + // A CI runner that lost Docker must go red, not quietly report a + // pass with the check skipped. Local runs without Docker still skip, + // and so does macOS CI, whose runners have no Docker by design. + if os.Getenv("CI") != "" && runtime.GOOS != "darwin" { + Fail("testcontainers requires Docker and CI is set: " + err.Error()) + } + Skip("testcontainers requires Docker: " + err.Error()) + } + c, err := messaging.New(url) + if err != nil { + Fail("connecting to the test NATS server: " + err.Error()) + } + return c, c.Close + }) +}) diff --git a/core/services/messaging/client_validation_test.go b/core/services/messaging/client_validation_test.go new file mode 100644 index 000000000..ff7f611d1 --- /dev/null +++ b/core/services/messaging/client_validation_test.go @@ -0,0 +1,56 @@ +package messaging_test + +import ( + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/messaging" +) + +// The live client checks a subject before it touches the connection. A client +// with no connection at all proves the order: if any method reached the +// connection first it would panic or fail with a connection error instead of +// the subject sentinel, and no NATS server is needed to find out. +var _ = Describe("Client subject validation", func() { + var c *messaging.Client + + BeforeEach(func() { + c = &messaging.Client{} + }) + + calls := map[string]func(subject string) error{ + "Publish": func(s string) error { return c.Publish(s, map[string]string{"k": "v"}) }, + "Request": func(s string) error { + _, err := c.Request(s, []byte("{}"), time.Second) + return err + }, + "Subscribe": func(s string) error { + _, err := c.Subscribe(s, func([]byte) {}) + return err + }, + "QueueSubscribe": func(s string) error { + _, err := c.QueueSubscribe(s, "q", func([]byte) {}) + return err + }, + "SubscribeReply": func(s string) error { + _, err := c.SubscribeReply(s, func([]byte, func([]byte)) {}) + return err + }, + "QueueSubscribeReply": func(s string) error { + _, err := c.QueueSubscribeReply(s, "q", func([]byte, func([]byte)) {}) + return err + }, + } + + for name, call := range calls { + It(name+" refuses an unserved root before using the connection", func() { + Expect(call("bogus.thing")).To(MatchError(messaging.ErrUnservedSubject)) + }) + + It(name+" refuses a multi-token wildcard before using the connection", func() { + Expect(call("jobs.>")).To(MatchError(messaging.ErrUnsupportedWildcard)) + }) + } +}) diff --git a/core/services/messaging/export_test.go b/core/services/messaging/export_test.go new file mode 100644 index 000000000..8c6806752 --- /dev/null +++ b/core/services/messaging/export_test.go @@ -0,0 +1,5 @@ +package messaging + +// NATSRouteForTest exposes the route table so the external specs can pin the +// queue group per kind, which no producer-side behaviour reveals. +var NATSRouteForTest = natsRoute diff --git a/core/services/messaging/interfaces.go b/core/services/messaging/interfaces.go index 863b1d66e..a73b343a8 100644 --- a/core/services/messaging/interfaces.go +++ b/core/services/messaging/interfaces.go @@ -12,12 +12,20 @@ type Subscription interface { Unsubscribe() error } -// MessagingClient is the full interface for NATS messaging operations. -// Consumers should depend on this interface rather than the concrete Client -// for testability. -type MessagingClient interface { +// Broadcaster is the fan-out surface: a publish reaches every subscriber on +// every replica. Consumers that only publish and subscribe depend on this +// instead of the wide client, so a second carrier only has to honour two +// methods. Delivery is at-most-once, and a consumer must not read silence as +// evidence about a node. +type Broadcaster interface { Publisher Subscribe(subject string, handler func([]byte)) (Subscription, error) +} + +// MessagingClient is the full NATS surface: fan-out plus queue groups and +// request/reply. Only the code that owns a queue or a control request needs it. +type MessagingClient interface { + Broadcaster QueueSubscribe(subject, queue string, handler func([]byte)) (Subscription, error) QueueSubscribeReply(subject, queue string, handler func(data []byte, reply func([]byte))) (Subscription, error) SubscribeReply(subject string, handler func(data []byte, reply func([]byte))) (Subscription, error) diff --git a/core/services/messaging/interfaces_test.go b/core/services/messaging/interfaces_test.go new file mode 100644 index 000000000..a04cd2ce7 --- /dev/null +++ b/core/services/messaging/interfaces_test.go @@ -0,0 +1,14 @@ +package messaging_test + +import ( + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/testutil" +) + +// Compile-time conformance. A wide client is still a Broadcaster, and so is the +// in-memory double every cross-replica spec shares. +var ( + _ messaging.Broadcaster = messaging.MessagingClient(nil) + _ messaging.MessagingClient = (*messaging.Client)(nil) + _ messaging.Broadcaster = (*testutil.FakeBus)(nil) +) diff --git a/core/services/messaging/messagingtest/broadcaster.go b/core/services/messaging/messagingtest/broadcaster.go new file mode 100644 index 000000000..b0d7f37cb --- /dev/null +++ b/core/services/messaging/messagingtest/broadcaster.go @@ -0,0 +1,154 @@ +// Package messagingtest holds the conformance suite every fan-out carrier must +// pass. A carrier is run against it in its own package, so a behaviour one +// carrier has and another lacks is a red spec, not a production surprise. +package messagingtest + +import ( + "encoding/json" + "errors" + "strings" + "sync" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/messaging" +) + +// Factory returns a ready carrier and a cleanup for it. +type Factory func() (messaging.Broadcaster, func()) + +type blob struct { + Data string `json:"data"` +} + +// collector gathers payloads from a handler, safely across goroutines. +type collector struct { + mu sync.Mutex + msgs [][]byte +} + +func (c *collector) handler(b []byte) { + c.mu.Lock() + defer c.mu.Unlock() + c.msgs = append(c.msgs, append([]byte(nil), b...)) +} + +func (c *collector) count() int { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.msgs) +} + +func (c *collector) first() []byte { + c.mu.Lock() + defer c.mu.Unlock() + return c.msgs[0] +} + +// RunBroadcasterConformance registers the suite. Call it from a Describe-level +// position in the carrier's own test package. +func RunBroadcasterConformance(newBus Factory) { + Describe("Broadcaster conformance", func() { + var ( + bus messaging.Broadcaster + cleanup func() + ) + + BeforeEach(func() { bus, cleanup = newBus() }) + // A factory that skips (no Docker for the NATS run) returns before + // setting cleanup, so the guard keeps a skip from turning into a panic. + AfterEach(func() { + if cleanup != nil { + cleanup() + } + }) + + It("delivers a published message to every subscriber", func() { + a, b := &collector{}, &collector{} + _, err := bus.Subscribe("jobs.j1.progress", a.handler) + Expect(err).ToNot(HaveOccurred()) + _, err = bus.Subscribe("jobs.j1.progress", b.handler) + Expect(err).ToNot(HaveOccurred()) + + Expect(bus.Publish("jobs.j1.progress", blob{Data: "x"})).To(Succeed()) + + Eventually(a.count, 5*time.Second).Should(Equal(1)) + Eventually(b.count, 5*time.Second).Should(Equal(1)) + }) + + It("matches a single-token wildcard and nothing wider", func() { + c := &collector{} + _, err := bus.Subscribe("jobs.*.cancel", c.handler) + Expect(err).ToNot(HaveOccurred()) + + Expect(bus.Publish("jobs.abc.result", blob{Data: "no"})).To(Succeed()) + Expect(bus.Publish("jobs.abc.cancel", blob{Data: "yes"})).To(Succeed()) + + Eventually(c.count, 5*time.Second).Should(Equal(1)) + Consistently(c.count, 300*time.Millisecond).Should(Equal(1)) + }) + + It("stops delivering after Unsubscribe", func() { + c := &collector{} + sub, err := bus.Subscribe("gallery.op1.progress", c.handler) + Expect(err).ToNot(HaveOccurred()) + Expect(bus.Publish("gallery.op1.progress", blob{Data: "1"})).To(Succeed()) + Eventually(c.count, 5*time.Second).Should(Equal(1)) + + Expect(sub.Unsubscribe()).To(Succeed()) + Expect(bus.Publish("gallery.op1.progress", blob{Data: "2"})).To(Succeed()) + Consistently(c.count, 300*time.Millisecond).Should(Equal(1)) + }) + + It("Unsubscribe removes only its own subscription", func() { + first, second := &collector{}, &collector{} + _, err := bus.Subscribe("gallery.op2.progress", first.handler) + Expect(err).ToNot(HaveOccurred()) + subSecond, err := bus.Subscribe("gallery.op2.progress", second.handler) + Expect(err).ToNot(HaveOccurred()) + + // Dropping the first one is the case a removal keyed on the subject + // gets right by accident, so drop the second and check the first + // survives. + Expect(subSecond.Unsubscribe()).To(Succeed()) + Expect(bus.Publish("gallery.op2.progress", blob{Data: "x"})).To(Succeed()) + + Eventually(first.count, 5*time.Second).Should(Equal(1)) + Consistently(second.count, 300*time.Millisecond).Should(Equal(0)) + }) + + It("refuses a subject outside the served roots on publish and subscribe", func() { + err := bus.Publish("bogus.thing", blob{Data: "x"}) + Expect(errors.Is(err, messaging.ErrUnservedSubject)).To(BeTrue(), "publish: %v", err) + + _, err = bus.Subscribe("bogus.thing", func([]byte) {}) + Expect(errors.Is(err, messaging.ErrUnservedSubject)).To(BeTrue(), "subscribe: %v", err) + }) + + It("refuses a multi-token wildcard subscription", func() { + _, err := bus.Subscribe("jobs.>", func([]byte) {}) + Expect(errors.Is(err, messaging.ErrUnsupportedWildcard)).To(BeTrue(), "got %v", err) + }) + + DescribeTable("delivers payloads past the PostgreSQL notify cap intact", + func(size int) { + c := &collector{} + _, err := bus.Subscribe("cache.invalidate.models", c.handler) + Expect(err).ToNot(HaveOccurred()) + + want := strings.Repeat("a", size) + Expect(bus.Publish("cache.invalidate.models", blob{Data: want})).To(Succeed()) + + Eventually(c.count, 10*time.Second).Should(Equal(1)) + var got blob + Expect(json.Unmarshal(c.first(), &got)).To(Succeed()) + Expect(got.Data).To(HaveLen(size)) + Expect(got.Data).To(Equal(want)) + }, + Entry("just under the notify cap", 7900), + Entry("over the notify cap", 64*1024), + ) + }) +} diff --git a/core/services/messaging/subject_rules.go b/core/services/messaging/subject_rules.go new file mode 100644 index 000000000..da0193bfe --- /dev/null +++ b/core/services/messaging/subject_rules.go @@ -0,0 +1,64 @@ +package messaging + +import ( + "errors" + "fmt" + "strings" +) + +// Every carrier serves the same closed set of subject roots. A subject outside +// it is refused at publish and at subscribe instead of being carried, because a +// subject that one carrier accepts and another drops is a message that is +// delivered to nobody, with no error anywhere. NATS would accept it, so the rule +// has to be stated here, once, rather than left to whichever carrier is in use. +var ( + // broadcastRoots carry fan-out and the competing-consumer subjects. + broadcastRoots = map[string]struct{}{ + "jobs": {}, "agent": {}, "gallery": {}, "cache": {}, + "staging": {}, "prefixcache": {}, "responses": {}, "state": {}, + "finetune": {}, + } + // controlRoots carry request/reply to one node or one agent worker. + controlRoots = map[string]struct{}{ + "nodes": {}, "mcp": {}, + } +) + +// ErrUnservedSubject is the class every root refusal belongs to, so a caller can +// tell "this carrier does not serve that family" from a transport failure +// without matching on strings. +var ErrUnservedSubject = errors.New("messaging: subject root is not served") + +// ErrUnsupportedWildcard reports a wildcard other than a whole single token. +// Only `*` standing alone in a non-root position is part of the contract; `>` is +// not, because not every carrier can honour it. +var ErrUnsupportedWildcard = errors.New("messaging: unsupported wildcard in subject") + +// ValidateSubject reports whether a subject, or a subscription filter, is one +// every carrier serves. +func ValidateSubject(subject string) error { + if subject == "" { + return fmt.Errorf("%w: empty subject", ErrUnservedSubject) + } + tokens := strings.Split(subject, ".") + for i, tok := range tokens { + switch { + case tok == "": + return fmt.Errorf("%w: %q has an empty token", ErrUnservedSubject, subject) + case strings.Contains(tok, ">"): + return fmt.Errorf("%w: %q", ErrUnsupportedWildcard, subject) + case strings.Contains(tok, "*") && tok != "*": + return fmt.Errorf("%w: %q", ErrUnsupportedWildcard, subject) + case tok == "*" && i == 0: + return fmt.Errorf("%w: %q has a wildcard root", ErrUnsupportedWildcard, subject) + } + } + root := tokens[0] + if _, ok := broadcastRoots[root]; ok { + return nil + } + if _, ok := controlRoots[root]; ok { + return nil + } + return fmt.Errorf("%w: %q", ErrUnservedSubject, subject) +} diff --git a/core/services/messaging/subject_rules_test.go b/core/services/messaging/subject_rules_test.go new file mode 100644 index 000000000..09549b152 --- /dev/null +++ b/core/services/messaging/subject_rules_test.go @@ -0,0 +1,100 @@ +package messaging_test + +import ( + "errors" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/messaging" +) + +var _ = Describe("Subject rules", func() { + DescribeTable("accepts served subjects", + func(subject string) { + Expect(messaging.ValidateSubject(subject)).To(Succeed()) + }, + Entry("job queue", "jobs.new"), + Entry("single-token wildcard", "jobs.*.cancel"), + Entry("two wildcards", "agent.*.events.*"), + Entry("control subject", "nodes.abc.backend.install"), + Entry("mcp request", "mcp.tools.execute"), + Entry("finetune progress", "finetune.job1.progress"), + ) + + DescribeTable("refuses a root that is not served", + func(subject string) { + err := messaging.ValidateSubject(subject) + Expect(errors.Is(err, messaging.ErrUnservedSubject)).To(BeTrue(), "got %v", err) + Expect(err.Error()).To(ContainSubstring(subject)) + }, + Entry("unknown root", "bogus.thing"), + Entry("empty token", "jobs..x"), + Entry("bare root that is unknown", "telemetry"), + ) + + It("refuses the empty subject", func() { + Expect(errors.Is(messaging.ValidateSubject(""), messaging.ErrUnservedSubject)).To(BeTrue()) + }) + + DescribeTable("refuses wildcards other than a whole single token", + func(subject string) { + Expect(errors.Is(messaging.ValidateSubject(subject), messaging.ErrUnsupportedWildcard)).To(BeTrue()) + }, + Entry("full wildcard", "jobs.>"), + Entry("partial token", "jobs.a*.cancel"), + Entry("wildcard root", "*.new"), + ) + + DescribeTable("serves every root", + func(root string) { + Expect(messaging.ValidateSubject(root + ".x")).To(Succeed()) + }, + Entry("jobs", "jobs"), Entry("agent", "agent"), Entry("gallery", "gallery"), + Entry("cache", "cache"), Entry("staging", "staging"), Entry("prefixcache", "prefixcache"), + Entry("responses", "responses"), Entry("state", "state"), Entry("finetune", "finetune"), + Entry("nodes", "nodes"), Entry("mcp", "mcp"), + ) + + It("refuses an unknown root", func() { + Expect(errors.Is(messaging.ValidateSubject("bogus.x"), messaging.ErrUnservedSubject)).To(BeTrue()) + }) + + It("serves every subject the constructors in subjects.go build", func() { + const id = "11111111-2222-3333-4444-555555555555" + // When you add a subject constant or constructor, add it here too. + subjects := []string{ + messaging.SubjectJobsNew, messaging.SubjectMCPCIJobsNew, messaging.SubjectAgentExecute, + messaging.SubjectMCPToolExecute, messaging.SubjectMCPDiscovery, + messaging.SubjectGalleryOpStart, messaging.SubjectGalleryOpEnd, + messaging.SubjectCacheInvalidateSkills, messaging.SubjectCacheInvalidateModels, + messaging.SubjectCacheInvalidateBackends, + messaging.SubjectPrefixCacheObserve, messaging.SubjectPrefixCacheInvalidate, + messaging.SubjectPrefixCachePressure, messaging.SubjectPrefixCacheResidency, + messaging.SubjectJobResultWildcard, + messaging.SubjectJobProgressWildcard, messaging.SubjectAgentCancelWildcard, + messaging.SubjectGalleryCancelWildcard, messaging.SubjectGalleryProgressWildcard, + messaging.SubjectResponseCancelWildcard, + messaging.SubjectAgentEvents("agent1", "user1"), + messaging.SubjectJobProgress(id), messaging.SubjectJobResult(id), + messaging.SubjectFineTuneProgress(id), messaging.SubjectGalleryProgress(id), + messaging.SubjectStagingProgress(id), + messaging.SubjectJobCancel(id), messaging.SubjectAgentCancel(id), + messaging.SubjectFineTuneCancel(id), messaging.SubjectGalleryCancel(id), + messaging.SubjectResponseCancel(id), + 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.SubjectNodeModelUnload(id), messaging.SubjectNodeModelDelete(id), + messaging.SubjectNodeModelsRunning(id), messaging.SubjectNodeStop(id), + messaging.SubjectNodeFilesEnsure(id), messaging.SubjectNodeFilesStage(id), + messaging.SubjectNodeFilesRelease(id), messaging.SubjectNodeFilesTemp(id), + messaging.SubjectNodeFilesListDir(id), + messaging.SubjectNodeBackendInstallProgress(id, "op1"), + } + for _, s := range subjects { + Expect(messaging.ValidateSubject(s)).To(Succeed(), "subject %q", s) + } + }) +}) diff --git a/core/services/messaging/subjects.go b/core/services/messaging/subjects.go index dc3a435d9..89fab4b6b 100644 --- a/core/services/messaging/subjects.go +++ b/core/services/messaging/subjects.go @@ -101,7 +101,6 @@ const ( // Wildcard subjects for NATS subscriptions that match all IDs. const ( - SubjectJobCancelWildcard = "jobs.*.cancel" SubjectJobResultWildcard = "jobs.*.result" SubjectJobProgressWildcard = "jobs.*.progress" SubjectAgentCancelWildcard = "agent.*.cancel" @@ -160,44 +159,6 @@ func SubjectNodeBackendInstall(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".backend.install" } -// BackendInstallRequest is the payload for a backend.install NATS request. -type BackendInstallRequest struct { - Backend string `json:"backend"` - ModelID string `json:"model_id,omitempty"` - BackendGalleries string `json:"backend_galleries,omitempty"` - // URI is set for external installs (OCI image, URL, or path). When non-empty - // the worker routes to InstallExternalBackend instead of the gallery lookup. - URI string `json:"uri,omitempty"` - Name string `json:"name,omitempty"` - Alias string `json:"alias,omitempty"` - // ReplicaIndex selects which slot on the worker this load occupies, so two - // concurrent backend.install requests for the same model land on distinct - // gRPC processes and ports. Workers older than this field treat it as 0 - // (single-replica behavior — no collision because the controller never - // asks for replica > 0 on a node whose MaxReplicasPerModel is 1). - ReplicaIndex int32 `json:"replica_index,omitempty"` - // Force is retained on the wire only for backward compatibility with - // pre-2026-05-08 masters that did not know about backend.upgrade. New - // callers MUST send to SubjectNodeBackendUpgrade instead. Workers continue - // to honor Force=true here so a rolling update with new master + old - // worker still works (the master's install fallback path also uses this - // when backend.upgrade returns nats.ErrNoResponders). - Force bool `json:"force,omitempty"` - // OpID identifies the admin-side operation. When non-empty the worker - // publishes BackendInstallProgressEvent values to - // SubjectNodeBackendInstallProgress(nodeID, OpID) while the install is - // 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"` -} - -// BackendInstallReply is the response from a backend.install NATS request. -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"` -} - // SubjectNodeBackendUpgrade tells a worker node to force-reinstall a backend // from the gallery, stop every running process for that backend, and restart. // Uses NATS request-reply with a long deadline (gallery image pulls can take @@ -208,116 +169,19 @@ func SubjectNodeBackendUpgrade(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".backend.upgrade" } -// BackendUpgradeRequest is the payload for a backend.upgrade NATS request. -// It is intentionally a strict subset of BackendInstallRequest — there is no -// Force field because the upgrade subject IS the force semantics; no ModelID -// because upgrade is backend-scoped (it stops every replica using the binary -// before re-installing). Per-replica restart happens on the next routine load. -type BackendUpgradeRequest struct { - Backend string `json:"backend"` - BackendGalleries string `json:"backend_galleries,omitempty"` - URI string `json:"uri,omitempty"` - Name string `json:"name,omitempty"` - Alias string `json:"alias,omitempty"` - // ReplicaIndex is informational — upgrade stops all replicas regardless, - // but the field lets future per-replica metadata (e.g. progress reporting - // scoped to a slot) ride the same wire without a v3 type. - ReplicaIndex int32 `json:"replica_index,omitempty"` - // OpID identifies the admin-side operation. When non-empty the worker - // publishes BackendInstallProgressEvent values to - // SubjectNodeBackendInstallProgress(nodeID, OpID) while the force-reinstall - // runs, so the master can stream per-node progress for upgrades exactly as - // it already does for installs (an upgrade IS a force-reinstall, so the - // install-progress subject is reused rather than minting a new one — no new - // NATS permission or rolling-update compat surface). Empty on legacy callers. - OpID string `json:"op_id,omitempty"` -} - -// BackendUpgradeReply mirrors BackendInstallReply minus Address — upgrade does -// not start a process, so there is no port to advertise. The subsequent -// routine load will re-bind via backend.install and learn the new address. -type BackendUpgradeReply struct { - Success bool `json:"success"` - Error string `json:"error,omitempty"` - - // StoppedProcessKeys / ReportsStoppedProcesses carry the same - // stale-row-invalidation contract as on BackendDeleteReply; an upgrade - // force-stops every process using the binary and starts none back up, so it - // recycles ports exactly the way a delete does. See that type for why the - // boolean is not redundant with an empty list. - StoppedProcessKeys []string `json:"stopped_process_keys,omitempty"` - ReportsStoppedProcesses bool `json:"reports_stopped_processes,omitempty"` -} - // SubjectNodeBackendList queries a worker node for its installed backends. // Uses NATS request-reply. func SubjectNodeBackendList(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".backend.list" } -// BackendListRequest is the payload for a backend.list NATS request. -type BackendListRequest struct{} - -// BackendListReply is the response from a backend.list NATS request. -type BackendListReply struct { - Backends []NodeBackendInfo `json:"backends"` - Error string `json:"error,omitempty"` -} - -// NodeBackendInfo describes a backend installed on a worker node. -type NodeBackendInfo struct { - Name string `json:"name"` - IsSystem bool `json:"is_system"` - IsMeta bool `json:"is_meta"` - InstalledAt string `json:"installed_at,omitempty"` - GalleryURL string `json:"gallery_url,omitempty"` - // Version, URI and Digest enable cluster-wide upgrade detection — - // without them, the frontend cannot tell whether the installed OCI - // image matches the gallery entry, and upgrades silently never surface. - Version string `json:"version,omitempty"` - URI string `json:"uri,omitempty"` - Digest string `json:"digest,omitempty"` -} - -// BackendStopRequest controls worker-side process shutdown. Force skips the -// best-effort Free RPC so a backend stuck serving a request can still be -// terminated by the watchdog. -type BackendStopRequest struct { - Backend string `json:"backend"` - Force bool `json:"force,omitempty"` -} - -// BackendStopReply is the worker's answer to a backend.stop request. -// -// backend.stop had no reply until this type existed. The controller published -// and returned success as soon as the local publish succeeded, so a stop that -// killed nothing, and a stop that failed outright, both looked identical to a -// stop that worked. An operator calling the unload endpoint got HTTP 200 while -// the backend kept running and holding its VRAM. -type BackendStopReply struct { - Success bool `json:"success"` - Error string `json:"error,omitempty"` - - // StoppedProcessKeys names every `modelID#replica` process the worker - // terminated while serving this request. - StoppedProcessKeys []string `json:"stopped_process_keys,omitempty"` - - // ReportsStoppedProcesses distinguishes "this worker enumerates what it - // stopped and stopped nothing" from "this worker predates the field", the - // same way BackendDeleteReply does. Both send an empty list and only the - // first is authoritative, so a controller that cannot tell them apart would - // read silence as a completed stop — the exact conclusion this reply exists - // to prevent. - ReportsStoppedProcesses bool `json:"reports_stopped_processes,omitempty"` -} - // SubjectNodeBackendStop tells a worker node to stop its gRPC backend process. // Equivalent to the local deleteProcess(). The node will: // 1. Best-effort bounded Free() via gRPC (unless Force is true) // 2. Kill the backend process // 3. Can be restarted via another backend.start event. // -// Request-reply, answered with a BackendStopReply. A worker that predates that +// Request-reply, answered with a workerctl.BackendStopReply. A worker that predates that // reply never answers, so the controller must treat a timeout as "unconfirmed" // rather than "failed" — see RemoteUnloaderAdapter.stopBackend. func SubjectNodeBackendStop(nodeID string) string { @@ -330,92 +194,24 @@ func SubjectNodeModelStop(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".model.stop" } -type ModelStopRequest struct { - ModelName string `json:"model_name"` - ProcessKey string `json:"process_key"` - ExpectedAddress string `json:"expected_address"` - Force bool `json:"force,omitempty"` - ConfigRevision string `json:"config_revision,omitempty"` -} - -type ModelStopReply struct { - Matched bool `json:"matched"` - Freed bool `json:"freed"` - Terminated bool `json:"terminated"` - ProcessKey string `json:"process_key"` - Address string `json:"address,omitempty"` - Error string `json:"error,omitempty"` -} - // SubjectNodeBackendDelete tells a worker node to delete a backend (stop + remove files). // Uses NATS request-reply. func SubjectNodeBackendDelete(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".backend.delete" } -// BackendDeleteRequest is the payload for a backend.delete NATS request. -type BackendDeleteRequest struct { - Backend string `json:"backend"` -} - -// BackendDeleteReply is the response from a backend.delete NATS request. -type BackendDeleteReply struct { - Success bool `json:"success"` - Error string `json:"error,omitempty"` - - // StoppedProcessKeys names every `modelID#replica` process the worker - // terminated while serving this delete. Stopping a process returns its gRPC - // port to the worker's allocator, so any NodeModel row still pointing at - // that address becomes a live misroute the moment an unrelated backend - // binds the recycled port: probeHealth verifies liveness, not identity, so - // the request is served by the wrong backend rather than failing. The - // controller uses these keys to drop the rows eagerly. - StoppedProcessKeys []string `json:"stopped_process_keys,omitempty"` - - // ReportsStoppedProcesses distinguishes "this worker enumerates what it - // stopped and stopped nothing" from "this worker predates the field". Both - // send an empty list, and only the first is authoritative. Without this - // flag a controller cannot tell them apart and would eventually be tempted - // to read silence as a completed cleanup, which is precisely the wrong - // conclusion against an older worker. - ReportsStoppedProcesses bool `json:"reports_stopped_processes,omitempty"` -} - // SubjectNodeModelUnload tells a worker node to unload a model (gRPC Free) without killing the backend. // Uses NATS request-reply. func SubjectNodeModelUnload(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".model.unload" } -// ModelUnloadRequest is the payload for a model.unload NATS request. -type ModelUnloadRequest struct { - ModelName string `json:"model_name"` - Address string `json:"address,omitempty"` // gRPC address of the backend process to unload from -} - -// ModelUnloadReply is the response from a model.unload NATS request. -type ModelUnloadReply struct { - Success bool `json:"success"` - Error string `json:"error,omitempty"` -} - // SubjectNodeModelDelete tells a worker node to delete model files from disk. // Uses NATS request-reply. func SubjectNodeModelDelete(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".model.delete" } -// ModelDeleteRequest is the payload for a model.delete NATS request. -type ModelDeleteRequest struct { - ModelName string `json:"model_name"` -} - -// ModelDeleteReply is the response from a model.delete NATS request. -type ModelDeleteReply struct { - Success bool `json:"success"` - Error string `json:"error,omitempty"` -} - // SubjectNodeModelsRunning asks a worker node which model backend processes it // currently has running. Uses NATS request-reply. // @@ -428,24 +224,6 @@ func SubjectNodeModelsRunning(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".models.running" } -// ModelsRunningRequest is the payload for a models.running NATS request. -type ModelsRunningRequest struct{} - -// ModelsRunningReply is the response from a models.running NATS request. -type ModelsRunningReply struct { - Models []RunningModelInfo `json:"models"` - Error string `json:"error,omitempty"` -} - -// RunningModelInfo identifies one live backend process on a worker. The triple -// is isomorphic to a controller NodeModel row's (model_name, replica_index, -// address), which is what lets the reconciler diff the two directly. -type RunningModelInfo struct { - ModelID string `json:"model_id"` - ReplicaIndex int `json:"replica_index"` - Address string `json:"address,omitempty"` -} - // SubjectNodeStop tells a serve-backend node to shut down entirely // (deregister + exit). The node will not restart the backend process. func SubjectNodeStop(nodeID string) string { @@ -456,31 +234,31 @@ func SubjectNodeStop(nodeID string) string { // These subjects use request-reply for synchronous file operations. // SubjectNodeFilesEnsure tells a serve-backend node to download an S3 key to its local cache. -// Reply: {local_path, error} +// Reply: workerctl.FileEnsureReply func SubjectNodeFilesEnsure(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".files.ensure" } // SubjectNodeFilesStage tells a serve-backend node to upload a local file to S3. -// Reply: {key, error} +// Reply: workerctl.FileStageReply func SubjectNodeFilesStage(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".files.stage" } // SubjectNodeFilesRelease tells a serve-backend node to evict one request's ephemeral cache keys. -// Reply: {error} +// Reply: workerctl.FileReleaseReply func SubjectNodeFilesRelease(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".files.release" } // SubjectNodeFilesTemp tells a serve-backend node to allocate a temp file. -// Reply: {local_path, error} +// Reply: workerctl.FileTempReply func SubjectNodeFilesTemp(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".files.temp" } // SubjectNodeFilesListDir tells a serve-backend node to list files in a directory. -// Reply: {files: [...], error} +// Reply: workerctl.FileListDirReply func SubjectNodeFilesListDir(nodeID string) string { return subjectNodePrefix + sanitizeSubjectToken(nodeID) + ".files.listdir" } diff --git a/core/services/messaging/subjects_upgrade_test.go b/core/services/messaging/subjects_upgrade_test.go index e60369cfc..ce67b08d2 100644 --- a/core/services/messaging/subjects_upgrade_test.go +++ b/core/services/messaging/subjects_upgrade_test.go @@ -18,15 +18,3 @@ var _ = Describe("SubjectNodeBackendUpgrade", func() { To(Equal("nodes.a-b-c.backend.upgrade")) }) }) - -var _ = Describe("BackendUpgradeRequest", func() { - It("carries backend name, galleries JSON, and replica index", func() { - req := messaging.BackendUpgradeRequest{ - Backend: "llama-cpp", - BackendGalleries: `[{"name":"x"}]`, - ReplicaIndex: 2, - } - Expect(req.Backend).To(Equal("llama-cpp")) - Expect(req.ReplicaIndex).To(BeEquivalentTo(2)) - }) -}) diff --git a/core/services/messaging/workqueue.go b/core/services/messaging/workqueue.go new file mode 100644 index 000000000..898fc894b --- /dev/null +++ b/core/services/messaging/workqueue.go @@ -0,0 +1,45 @@ +package messaging + +import "context" + +// WorkKind names a unit of competing-consumer work. The values match the claim +// kinds of the self-hosted carrier so the two map one to one. +type WorkKind string + +const ( + WorkTask WorkKind = "task" // NATS: jobs.new, group "workers" + WorkMCPCI WorkKind = "mcp-ci" // NATS: jobs.mcp-ci.new, group "workers" + WorkAgentRun WorkKind = "agent-run" // NATS: agent.execute, group "agent-workers" +) + +// WorkQueue is the producer side of competing-consumer work: exactly one +// consumer of the kind is meant to run each payload. +// +// A nil error means the carrier accepted the payload, not that any consumer +// exists or will run it. The NATS carrier refuses a payload over the server's +// max_payload (1 MB by default); the Broadcaster's 7999-byte guarantee does not +// apply to the queue. +// +// The NATS carrier ignores ctx: its publish takes none and returns once the +// message is buffered, so there is nothing for a cancellation to interrupt. +type WorkQueue interface { + Enqueue(ctx context.Context, kind WorkKind, payload any) error +} + +// WorkHandler runs one unit of work to its conclusion and returns only then. +// events is where the work publishes its progress, results and agent events. A +// nil return means this worker ran the work (success or failure is reported on +// events); a non-nil return means it could not serve it. A payload that can +// never decode returns nil: the handler logs and drops it, because an error +// would only add a carrier warning and could make a redelivering carrier loop +// on it. Delivery count is carrier-defined: a handler must tolerate a repeat. +type WorkHandler func(ctx context.Context, payload []byte, events Publisher) error + +// WorkConsumer is the worker side. ctx is the parent of every handler call. +// maxInFlight bounds concurrent handler calls: 0 is unbounded, 1 is serial, +// and a negative value is unbounded like 0. Unsubscribe stops delivery and +// waits for in-flight handlers to return, so calling it from inside a handler +// deadlocks. +type WorkConsumer interface { + Consume(ctx context.Context, kind WorkKind, maxInFlight int, h WorkHandler) (Subscription, error) +} diff --git a/core/services/messaging/workqueue_nats.go b/core/services/messaging/workqueue_nats.go new file mode 100644 index 000000000..a4b58fe0c --- /dev/null +++ b/core/services/messaging/workqueue_nats.go @@ -0,0 +1,196 @@ +package messaging + +import ( + "context" + "fmt" + "sync" + + "github.com/mudler/LocalAI/pkg/concurrency" + "github.com/mudler/xlog" +) + +// natsRoute is the one place that maps a kind onto a NATS subject and queue +// group. jobs.new and jobs.mcp-ci.new share the "workers" group on purpose: +// changing a group changes which processes compete for a message. +func natsRoute(kind WorkKind) (subject, queue string, err error) { + switch kind { + case WorkTask: + return SubjectJobsNew, QueueWorkers, nil + case WorkMCPCI: + return SubjectMCPCIJobsNew, QueueWorkers, nil + case WorkAgentRun: + return SubjectAgentExecute, QueueAgentWorkers, nil + default: + return "", "", fmt.Errorf("unknown work kind %q", kind) + } +} + +type natsWorkQueue struct { + pub Publisher +} + +// NewNATSWorkQueue returns a WorkQueue that publishes on the kind's subject. +// The queue group is a consumer concern, so the producer only needs Publish. +func NewNATSWorkQueue(pub Publisher) WorkQueue { + return &natsWorkQueue{pub: pub} +} + +func (q *natsWorkQueue) Enqueue(_ context.Context, kind WorkKind, payload any) error { + subject, _, err := natsRoute(kind) + if err != nil { + return err + } + // Publish marshals payload itself; handing it pre-encoded bytes would + // double encode. + return q.pub.Publish(subject, payload) +} + +type natsWorkRoutes struct { + agentSubject, agentQueue string + // agentQueueSet tells an explicitly empty queue from no override, because + // the two subscribe differently. + agentQueueSet bool +} + +// WorkRouteOption changes where a NATS WorkConsumer listens. +type WorkRouteOption func(*natsWorkRoutes) + +// WithAgentRunRoute moves the agent-run subject and queue group, which workers +// let operators set (LOCALAI_AGENT_SUBJECT, LOCALAI_AGENT_QUEUE). The other +// kinds have no such setting. +// +// An empty subject keeps the default subject. An empty queue is kept as given: +// it means no queue group, a plain subscription where every agent worker runs +// every agent run. That is what an explicitly empty LOCALAI_AGENT_QUEUE has +// always done, so it stays the operator's choice. +func WithAgentRunRoute(subject, queue string) WorkRouteOption { + return func(r *natsWorkRoutes) { + r.agentSubject, r.agentQueue, r.agentQueueSet = subject, queue, true + } +} + +type natsWorkConsumer struct { + c MessagingClient + routes natsWorkRoutes +} + +// NewNATSWorkConsumer returns a WorkConsumer that joins the kind's queue group. +// The handler's events publisher is c itself. +func NewNATSWorkConsumer(c MessagingClient, opts ...WorkRouteOption) WorkConsumer { + w := &natsWorkConsumer{c: c} + for _, o := range opts { + o(&w.routes) + } + return w +} + +func (w *natsWorkConsumer) route(kind WorkKind) (subject, queue string, err error) { + subject, queue, err = natsRoute(kind) + if err != nil || kind != WorkAgentRun { + return subject, queue, err + } + if w.routes.agentSubject != "" { + subject = w.routes.agentSubject + } + if w.routes.agentQueueSet { + queue = w.routes.agentQueue + } + return subject, queue, nil +} + +// natsWorkSubscription tracks in-flight handlers so Unsubscribe can wait for +// them. closed stops a delivery that got past NATS before the unsubscribe from +// calling wg.Add while Unsubscribe is already in wg.Wait. +type natsWorkSubscription struct { + sub Subscription + mu sync.Mutex + closed bool + wg sync.WaitGroup +} + +func (s *natsWorkSubscription) begin() bool { + s.mu.Lock() + defer s.mu.Unlock() + if s.closed { + return false + } + s.wg.Add(1) + return true +} + +// Unsubscribe stops delivery first so no new handler starts, then waits for +// the running ones, as the agent dispatcher's Stop always has. +func (s *natsWorkSubscription) Unsubscribe() error { + err := s.sub.Unsubscribe() + s.mu.Lock() + s.closed = true + s.mu.Unlock() + s.wg.Wait() + return err +} + +// Consume keeps the two concurrency models the workers had before this seam. +// With maxInFlight 1 the handler runs on the NATS delivery goroutine: one unit +// at a time per worker, messages already handed to this subscription wait +// behind it, and a panic is not recovered. Any other value spawns a recovered +// goroutine per delivery; when bounded, the slot is taken on the delivery +// goroutine, so a full worker stops draining its subscription instead of +// piling up goroutines, and ctx cancellation releases a delivery that waits. +func (w *natsWorkConsumer) Consume(ctx context.Context, kind WorkKind, maxInFlight int, h WorkHandler) (Subscription, error) { + subject, queue, err := w.route(kind) + if err != nil { + return nil, err + } + ws := &natsWorkSubscription{} + run := func(payload []byte) { + if err := h(ctx, payload, w.c); err != nil { + xlog.Warn("Work handler could not serve a delivery", "kind", kind, "subject", subject, "error", err) + } + } + + var deliver func([]byte) + switch { + case maxInFlight == 1: + deliver = func(payload []byte) { + if !ws.begin() { + return + } + defer ws.wg.Done() + run(payload) + } + default: + var sem chan struct{} + if maxInFlight > 0 { + sem = make(chan struct{}, maxInFlight) + } + deliver = func(payload []byte) { + if sem != nil { + select { + case sem <- struct{}{}: + case <-ctx.Done(): + return + } + } + if !ws.begin() { + if sem != nil { + <-sem + } + return + } + concurrency.SafeGo(func() { + defer ws.wg.Done() + if sem != nil { + defer func() { <-sem }() + } + run(payload) + }) + } + } + + sub, err := w.c.QueueSubscribe(subject, queue, deliver) + if err != nil { + return nil, fmt.Errorf("subscribing to %s: %w", subject, err) + } + ws.sub = sub + return ws, nil +} diff --git a/core/services/messaging/workqueue_nats_test.go b/core/services/messaging/workqueue_nats_test.go new file mode 100644 index 000000000..53fb9e3e1 --- /dev/null +++ b/core/services/messaging/workqueue_nats_test.go @@ -0,0 +1,346 @@ +package messaging_test + +import ( + "context" + "encoding/json" + "fmt" + "sync" + "sync/atomic" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/testutil" +) + +var _ = Describe("NATS work queue", func() { + DescribeTable("routes each kind to its subject and marshals the payload once", + func(kind messaging.WorkKind, subject string) { + bus := testutil.NewFakeBus() + var got []byte + _, err := bus.Subscribe(subject, func(b []byte) { got = b }) + Expect(err).ToNot(HaveOccurred()) + + Expect(messaging.NewNATSWorkQueue(bus).Enqueue(context.Background(), kind, map[string]string{"id": "x"})).To(Succeed()) + + Expect(bus.PublishCount(subject)).To(Equal(1)) + var back map[string]string + Expect(json.Unmarshal(got, &back)).To(Succeed()) + Expect(back).To(Equal(map[string]string{"id": "x"})) + }, + Entry("task", messaging.WorkTask, "jobs.new"), + Entry("mcp ci", messaging.WorkMCPCI, "jobs.mcp-ci.new"), + Entry("agent run", messaging.WorkAgentRun, "agent.execute"), + ) + + DescribeTable("pins the subject and queue group per kind", + func(kind messaging.WorkKind, subject, queue string) { + gotSubject, gotQueue, err := messaging.NATSRouteForTest(kind) + Expect(err).ToNot(HaveOccurred()) + Expect(gotSubject).To(Equal(subject)) + Expect(gotQueue).To(Equal(queue)) + }, + Entry("task", messaging.WorkTask, "jobs.new", "workers"), + Entry("mcp ci shares the task group", messaging.WorkMCPCI, "jobs.mcp-ci.new", "workers"), + Entry("agent run", messaging.WorkAgentRun, "agent.execute", "agent-workers"), + ) + + It("refuses an unknown kind without publishing", func() { + bus := testutil.NewFakeBus() + err := messaging.NewNATSWorkQueue(bus).Enqueue(context.Background(), messaging.WorkKind("nope"), 1) + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("nope")) + for _, s := range []string{"jobs.new", "jobs.mcp-ci.new", "agent.execute"} { + Expect(bus.PublishCount(s)).To(BeZero()) + } + }) + + It("publishes even when ctx is already cancelled, because the NATS publish takes no ctx", func() { + bus := testutil.NewFakeBus() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + Expect(messaging.NewNATSWorkQueue(bus).Enqueue(ctx, messaging.WorkTask, map[string]string{"id": "x"})).To(Succeed()) + Expect(bus.PublishCount("jobs.new")).To(Equal(1)) + }) +}) + +// deliverInOrder publishes each payload on subject one after the other from a +// single goroutine, which is how NATS feeds one subscription: the next message +// reaches the callback only once the previous callback has returned. FakeBus +// delivers synchronously inside Publish, so this reproduces that goroutine. +func deliverInOrder(bus *testutil.FakeBus, subject string, payloads ...string) <-chan struct{} { + done := make(chan struct{}) + go func() { + defer GinkgoRecover() + defer close(done) + for _, p := range payloads { + Expect(bus.Publish(subject, p)).To(Succeed()) + } + }() + return done +} + +// gatedHandler reports each start on started and holds every call until its +// payload's gate is closed. +type gatedHandler struct { + mu sync.Mutex + gates map[string]chan struct{} + started chan string + running atomic.Int32 + peak atomic.Int32 +} + +func newGatedHandler(payloads ...string) *gatedHandler { + g := &gatedHandler{gates: map[string]chan struct{}{}, started: make(chan string, 16)} + for _, p := range payloads { + g.gates[fmt.Sprintf("%q", p)] = make(chan struct{}) + } + return g +} + +func (g *gatedHandler) release(p string) { close(g.gates[fmt.Sprintf("%q", p)]) } + +func (g *gatedHandler) handle(_ context.Context, payload []byte, _ messaging.Publisher) error { + n := g.running.Add(1) + defer g.running.Add(-1) + for { + p := g.peak.Load() + if n <= p || g.peak.CompareAndSwap(p, n) { + break + } + } + g.mu.Lock() + gate := g.gates[string(payload)] + g.mu.Unlock() + g.started <- string(payload) + <-gate + return nil +} + +var _ = Describe("NATS work consumer", func() { + var ( + bus *testutil.FakeBus + wc messaging.WorkConsumer + ) + + BeforeEach(func() { + bus = testutil.NewFakeBus() + wc = messaging.NewNATSWorkConsumer(bus) + }) + + It("runs the handler inline with an in-flight limit of one, so the next delivery waits", func() { + g := newGatedHandler("a", "b") + sub, err := wc.Consume(context.Background(), messaging.WorkMCPCI, 1, g.handle) + Expect(err).ToNot(HaveOccurred()) + + done := deliverInOrder(bus, "jobs.mcp-ci.new", "a", "b") + Eventually(g.started).Should(Receive(Equal(`"a"`))) + Consistently(g.started, "100ms").ShouldNot(Receive()) + + g.release("a") + Eventually(g.started).Should(Receive(Equal(`"b"`))) + g.release("b") + Eventually(done).Should(BeClosed()) + Expect(g.peak.Load()).To(Equal(int32(1))) + Expect(sub.Unsubscribe()).To(Succeed()) + }) + + It("overlaps handlers with no in-flight limit", func() { + g := newGatedHandler("a", "b") + sub, err := wc.Consume(context.Background(), messaging.WorkAgentRun, 0, g.handle) + Expect(err).ToNot(HaveOccurred()) + + done := deliverInOrder(bus, "agent.execute", "a", "b") + Eventually(done).Should(BeClosed()) + Eventually(g.started).Should(Receive()) + Eventually(g.started).Should(Receive()) + Expect(g.running.Load()).To(Equal(int32(2))) + + g.release("a") + g.release("b") + Expect(sub.Unsubscribe()).To(Succeed()) + }) + + It("holds a delivery on the delivery goroutine until a slot frees with a limit of two", func() { + g := newGatedHandler("a", "b", "c") + sub, err := wc.Consume(context.Background(), messaging.WorkAgentRun, 2, g.handle) + Expect(err).ToNot(HaveOccurred()) + + done := deliverInOrder(bus, "agent.execute", "a", "b", "c") + Eventually(g.started).Should(Receive()) + Eventually(g.started).Should(Receive()) + Consistently(g.started, "100ms").ShouldNot(Receive()) + Expect(done).ToNot(BeClosed()) + + g.release("a") + Eventually(g.started).Should(Receive(Equal(`"c"`))) + Eventually(done).Should(BeClosed()) + g.release("b") + g.release("c") + Expect(sub.Unsubscribe()).To(Succeed()) + Expect(g.peak.Load()).To(Equal(int32(2))) + }) + + DescribeTable("Unsubscribe returns only after the in-flight handler returns", + func(maxInFlight int) { + g := newGatedHandler("a") + sub, err := wc.Consume(context.Background(), messaging.WorkAgentRun, maxInFlight, g.handle) + Expect(err).ToNot(HaveOccurred()) + + deliverInOrder(bus, "agent.execute", "a") + Eventually(g.started).Should(Receive()) + + unsubscribed := make(chan struct{}) + go func() { + defer GinkgoRecover() + defer close(unsubscribed) + Expect(sub.Unsubscribe()).To(Succeed()) + }() + Consistently(unsubscribed, "100ms").ShouldNot(BeClosed()) + + g.release("a") + Eventually(unsubscribed).Should(BeClosed()) + Expect(bus.Publish("agent.execute", "late")).To(Succeed()) + Consistently(g.started, "50ms").ShouldNot(Receive()) + }, + Entry("unbounded", 0), + Entry("inline", 1), + Entry("bounded", 2), + ) + + It("lets a cancelled ctx release a delivery that is waiting for a slot", func() { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + g := newGatedHandler("a", "b") + sub, err := wc.Consume(ctx, messaging.WorkAgentRun, 1+1, g.handle) + Expect(err).ToNot(HaveOccurred()) + + // Two slots, both held, so the third delivery waits for one. + done := deliverInOrder(bus, "agent.execute", "a", "b", "c") + Eventually(g.started).Should(Receive()) + Eventually(g.started).Should(Receive()) + Consistently(done, "100ms").ShouldNot(BeClosed()) + + cancel() + Eventually(done).Should(BeClosed()) + Consistently(g.started, "50ms").ShouldNot(Receive()) + + g.release("a") + g.release("b") + Expect(sub.Unsubscribe()).To(Succeed()) + }) + + DescribeTable("subscribes each kind on its default subject and queue group", + func(kind messaging.WorkKind, subject, queue string) { + sub, err := wc.Consume(context.Background(), kind, 0, func(context.Context, []byte, messaging.Publisher) error { return nil }) + Expect(err).ToNot(HaveOccurred()) + Expect(bus.QueueGroups()).To(Equal(map[string]string{subject: queue})) + Expect(sub.Unsubscribe()).To(Succeed()) + }, + Entry("task", messaging.WorkTask, "jobs.new", "workers"), + Entry("mcp ci shares the task group", messaging.WorkMCPCI, "jobs.mcp-ci.new", "workers"), + Entry("agent run", messaging.WorkAgentRun, "agent.execute", "agent-workers"), + ) + + DescribeTable("moves only the agent-run route when it is overridden", + func(kind messaging.WorkKind, subject, queue string) { + wc := messaging.NewNATSWorkConsumer(bus, messaging.WithAgentRunRoute("agent.tenant.x", "q2")) + sub, err := wc.Consume(context.Background(), kind, 0, func(context.Context, []byte, messaging.Publisher) error { return nil }) + Expect(err).ToNot(HaveOccurred()) + Expect(bus.QueueGroups()).To(Equal(map[string]string{subject: queue})) + Expect(sub.Unsubscribe()).To(Succeed()) + }, + Entry("task keeps its route", messaging.WorkTask, "jobs.new", "workers"), + Entry("mcp ci keeps its route", messaging.WorkMCPCI, "jobs.mcp-ci.new", "workers"), + Entry("agent run moves", messaging.WorkAgentRun, "agent.tenant.x", "q2"), + ) + + // LOCALAI_AGENT_QUEUE="" used to reach QueueSubscribe as given, which is a + // plain subscription where every agent worker runs every run. An empty + // subject has no such meaning, so it falls back to the default. + It("keeps an explicitly empty agent-run queue as a plain subscription", func() { + wc := messaging.NewNATSWorkConsumer(bus, messaging.WithAgentRunRoute("agent.tenant.x", "")) + sub, err := wc.Consume(context.Background(), messaging.WorkAgentRun, 0, func(context.Context, []byte, messaging.Publisher) error { return nil }) + Expect(err).ToNot(HaveOccurred()) + Expect(bus.QueueGroups()).To(Equal(map[string]string{"agent.tenant.x": ""})) + Expect(sub.Unsubscribe()).To(Succeed()) + }) + + It("falls back to the default agent-run subject when the subject is empty", func() { + wc := messaging.NewNATSWorkConsumer(bus, messaging.WithAgentRunRoute("", "q2")) + sub, err := wc.Consume(context.Background(), messaging.WorkAgentRun, 0, func(context.Context, []byte, messaging.Publisher) error { return nil }) + Expect(err).ToNot(HaveOccurred()) + Expect(bus.QueueGroups()).To(Equal(map[string]string{"agent.execute": "q2"})) + Expect(sub.Unsubscribe()).To(Succeed()) + }) + + It("refuses an agent-run override on an unserved subject", func() { + wc := messaging.NewNATSWorkConsumer(bus, messaging.WithAgentRunRoute("bogus.x", "q")) + _, err := wc.Consume(context.Background(), messaging.WorkAgentRun, 0, func(context.Context, []byte, messaging.Publisher) error { return nil }) + Expect(err).To(MatchError(messaging.ErrUnservedSubject)) + Expect(bus.QueueGroups()).To(BeEmpty()) + }) + + It("refuses an unknown kind", func() { + _, err := wc.Consume(context.Background(), messaging.WorkKind("nope"), 0, func(context.Context, []byte, messaging.Publisher) error { return nil }) + Expect(err).To(HaveOccurred()) + Expect(err.Error()).To(ContainSubstring("nope")) + }) + + It("hands the handler the raw payload, the consume ctx and the bus as its publisher", func() { + type ctxKey struct{} + ctx := context.WithValue(context.Background(), ctxKey{}, "parent") + got := make(chan []byte, 1) + var gotCtx context.Context + var gotEvents messaging.Publisher + sub, err := wc.Consume(ctx, messaging.WorkMCPCI, 1, func(hctx context.Context, payload []byte, events messaging.Publisher) error { + gotCtx, gotEvents = hctx, events + got <- payload + return nil + }) + Expect(err).ToNot(HaveOccurred()) + + // The bytes arrive as published, still encoded: decoding is the handler's job. + Expect(bus.Publish("jobs.mcp-ci.new", json.RawMessage(`[1,2]`))).To(Succeed()) + Eventually(got).Should(Receive(Equal([]byte(`[1,2]`)))) + Expect(gotCtx.Value(ctxKey{})).To(Equal("parent")) + Expect(gotEvents).To(BeIdenticalTo(bus)) + Expect(sub.Unsubscribe()).To(Succeed()) + }) + + It("recovers a handler panic when handlers are spawned", func() { + ran := make(chan string, 2) + sub, err := wc.Consume(context.Background(), messaging.WorkAgentRun, 0, func(_ context.Context, payload []byte, _ messaging.Publisher) error { + ran <- string(payload) + if string(payload) == `"boom"` { + panic("boom") + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + + done := deliverInOrder(bus, "agent.execute", "boom", "ok") + Eventually(done).Should(BeClosed()) + Eventually(ran).Should(Receive()) + Eventually(ran).Should(Receive()) + Expect(sub.Unsubscribe()).To(Succeed()) + }) + + It("does not recover a handler panic with an in-flight limit of one, as the inline consumer never did", func() { + sub, err := wc.Consume(context.Background(), messaging.WorkMCPCI, 1, func(context.Context, []byte, messaging.Publisher) error { + panic("boom") + }) + Expect(err).ToNot(HaveOccurred()) + + // The panic surfaces on the delivery goroutine, which on NATS is the + // client's own goroutine and so ends the process. Catch it there. + recovered := make(chan any, 1) + go func() { + defer func() { recovered <- recover() }() + _ = bus.Publish("jobs.mcp-ci.new", "x") + }() + Eventually(recovered).Should(Receive(Equal("boom"))) + Expect(sub.Unsubscribe()).To(Succeed()) + }) +}) diff --git a/core/services/nodes/agent_rpc_nats.go b/core/services/nodes/agent_rpc_nats.go new file mode 100644 index 000000000..c68426d32 --- /dev/null +++ b/core/services/nodes/agent_rpc_nats.go @@ -0,0 +1,132 @@ +package nodes + +import ( + "context" + "encoding/json" + "fmt" + "time" + + "github.com/mudler/LocalAI/core/config" + mcpremote "github.com/mudler/LocalAI/core/services/mcp" + "github.com/mudler/LocalAI/core/services/messaging" +) + +// NATSAgentControl sends the frontend's MCP verbs to the agent-worker queue +// group as NATS request/reply. +type NATSAgentControl struct{ bus messaging.MessagingClient } + +func NewNATSAgentControl(bus messaging.MessagingClient) *NATSAgentControl { + return &NATSAgentControl{bus: bus} +} + +func (a *NATSAgentControl) ExecuteMCPTool(ctx context.Context, req mcpremote.MCPToolRequest) (*mcpremote.MCPToolResponse, error) { + timeout, err := agentRequestTimeout(ctx, config.DefaultMCPToolTimeout) + if err != nil { + return nil, err + } + return controlRequestJSON[mcpremote.MCPToolRequest, mcpremote.MCPToolResponse](a.bus, messaging.SubjectMCPToolExecute, req, timeout) +} + +func (a *NATSAgentControl) DiscoverMCPTools(ctx context.Context, req mcpremote.MCPDiscoveryRequest) (*mcpremote.MCPDiscoveryResponse, error) { + timeout, err := agentRequestTimeout(ctx, config.DefaultMCPDiscoveryTimeout) + if err != nil { + return nil, err + } + return controlRequestJSON[mcpremote.MCPDiscoveryRequest, mcpremote.MCPDiscoveryResponse](a.bus, messaging.SubjectMCPDiscovery, req, timeout) +} + +// agentRequestTimeout reads only the deadline of ctx, never its cancellation: +// a chat client that disconnects must not abort a tool call already running on +// a worker, which is how these requests have always behaved. +func agentRequestTimeout(ctx context.Context, fallback time.Duration) (time.Duration, error) { + deadline, ok := ctx.Deadline() + if !ok { + return fallback, nil + } + timeout := time.Until(deadline) + if timeout <= 0 { + return 0, fmt.Errorf("agent request budget spent: %w", context.DeadlineExceeded) + } + return timeout, nil +} + +// NATSAgentRPCServer is the agent worker's end of NATSAgentControl, plus the +// node's backend stop listener. It holds no subscription handles: they live as +// long as the process and nothing unsubscribes them. +type NATSAgentRPCServer struct { + bus messaging.MessagingClient + nodeID string +} + +// NewNATSAgentRPCServer serves on bus for the agent worker registered as +// nodeID. The node id scopes only the backend stop subject: the MCP requests +// are shared by the agent-workers queue group, so any worker may answer them. +func NewNATSAgentRPCServer(bus messaging.MessagingClient, nodeID string) *NATSAgentRPCServer { + return &NATSAgentRPCServer{bus: bus, nodeID: nodeID} +} + +// ServeMCPTool answers tool requests in the agent-workers queue group. The +// handler runs inline on the delivery goroutine, so a worker serves one tool +// call at a time, and on context.Background because the handler sets its own +// budget: a call in flight when the worker is told to stop runs to that budget. +func (s *NATSAgentRPCServer) ServeMCPTool(h mcpremote.ToolHandler) error { + _, err := s.bus.QueueSubscribeReply(messaging.SubjectMCPToolExecute, messaging.QueueAgentWorkers, func(data []byte, reply func([]byte)) { + var req mcpremote.MCPToolRequest + if err := json.Unmarshal(data, &req); err != nil { + sendAgentReply(reply, mcpremote.MCPToolResponse{Error: fmt.Sprintf("unmarshal error: %v", err)}) + return + } + sendAgentReply(reply, h(context.Background(), req)) + }) + if err != nil { + return fmt.Errorf("serving mcp tool on %s: %w", messaging.SubjectMCPToolExecute, err) + } + return nil +} + +// ServeMCPDiscovery answers discovery requests like ServeMCPTool answers tool +// requests. Its own subscription lets a discovery overlap a tool call. +func (s *NATSAgentRPCServer) ServeMCPDiscovery(h mcpremote.DiscoveryHandler) error { + _, err := s.bus.QueueSubscribeReply(messaging.SubjectMCPDiscovery, messaging.QueueAgentWorkers, func(data []byte, reply func([]byte)) { + var req mcpremote.MCPDiscoveryRequest + if err := json.Unmarshal(data, &req); err != nil { + sendAgentReply(reply, mcpremote.MCPDiscoveryResponse{Error: fmt.Sprintf("unmarshal error: %v", err)}) + return + } + sendAgentReply(reply, h(context.Background(), req)) + }) + if err != nil { + return fmt.Errorf("serving mcp discovery on %s: %w", messaging.SubjectMCPDiscovery, err) + } + return nil +} + +// ServeBackendStop calls h with the backend named by each stop request sent to +// this node. It never replies: the frontend sends the stop as a request and +// reads the timeout from a node without a backend supervisor as success. A body +// it cannot decode is dropped, as nothing waits for an answer. +func (s *NATSAgentRPCServer) ServeBackendStop(h func(backend string)) error { + subject := messaging.SubjectNodeBackendStop(s.nodeID) + _, err := s.bus.Subscribe(subject, func(data []byte) { + // Only the backend name is read, so a change to the other fields of + // the stop request cannot make this listener drop it. + var req struct { + Backend string `json:"backend"` + } + if json.Unmarshal(data, &req) != nil { + return + } + h(req.Backend) + }) + if err != nil { + return fmt.Errorf("serving backend stop on %s: %w", subject, err) + } + return nil +} + +// sendAgentReply ignores the encoding error, as the agent worker always has: +// the response types hold only JSON-safe data. +func sendAgentReply(reply func([]byte), resp any) { + data, _ := json.Marshal(resp) + reply(data) +} diff --git a/core/services/nodes/agent_rpc_nats_test.go b/core/services/nodes/agent_rpc_nats_test.go new file mode 100644 index 000000000..919f0a891 --- /dev/null +++ b/core/services/nodes/agent_rpc_nats_test.go @@ -0,0 +1,388 @@ +package nodes + +import ( + "context" + "encoding/json" + "errors" + "sync/atomic" + "time" + + "github.com/nats-io/nats.go" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/config" + mcpremote "github.com/mudler/LocalAI/core/services/mcp" + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// The frontend reads ErrNoRoute as "no agent worker could be offered this", so +// a slow worker (timeout) or a worker's own error reply must never look like +// one. The timeout rules pin today's behaviour: only the deadline bounds the +// wait, and a client that goes away does not abort a tool call in flight. +var _ = Describe("NATS agent control", func() { + var ( + mc *scriptedMessagingClient + ac *NATSAgentControl + ) + + BeforeEach(func() { + mc = newScriptedMessagingClient() + ac = NewNATSAgentControl(mc) + }) + + lastTimeout := func() time.Duration { + mc.mu.Lock() + defer mc.mu.Unlock() + Expect(mc.calls).ToNot(BeEmpty()) + return mc.calls[len(mc.calls)-1].Timeout + } + + type verb struct { + subject string + fallback time.Duration + call func(ctx context.Context) (string, error) + errorReply any + } + + verbs := []struct { + name string + get func() verb + }{ + {"tool execution", func() verb { + return verb{ + subject: messaging.SubjectMCPToolExecute, + fallback: config.DefaultMCPToolTimeout, + errorReply: mcpremote.MCPToolResponse{Error: "tool 'x' not found"}, + call: func(ctx context.Context) (string, error) { + reply, err := ac.ExecuteMCPTool(ctx, mcpremote.MCPToolRequest{ModelName: "m", ToolName: "x"}) + if reply == nil { + return "", err + } + return reply.Error, err + }, + } + }}, + {"discovery", func() verb { + return verb{ + subject: messaging.SubjectMCPDiscovery, + fallback: config.DefaultMCPDiscoveryTimeout, + errorReply: mcpremote.MCPDiscoveryResponse{Error: "no MCP servers"}, + call: func(ctx context.Context) (string, error) { + reply, err := ac.DiscoverMCPTools(ctx, mcpremote.MCPDiscoveryRequest{ModelName: "m"}) + if reply == nil { + return "", err + } + return reply.Error, err + }, + } + }}, + } + + for _, v := range verbs { + Context(v.name, func() { + var vb verb + BeforeEach(func() { vb = v.get() }) + + It("reports no responders as ErrNoRoute without the carrier sentinel", func() { + mc.scriptNoResponders(vb.subject) + _, err := vb.call(context.Background()) + Expect(errors.Is(err, ErrNoRoute)).To(BeTrue(), "got %v", err) + Expect(errors.Is(err, nats.ErrNoResponders)).To(BeFalse()) + }) + + It("does not report a timeout as ErrNoRoute", func() { + mc.scriptErr(vb.subject, nats.ErrTimeout) + _, err := vb.call(context.Background()) + Expect(err).To(HaveOccurred()) + Expect(errors.Is(err, ErrNoRoute)).To(BeFalse()) + }) + + It("returns a reply carrying the worker's error with a nil error", func() { + mc.scriptReply(vb.subject, vb.errorReply) + workerErr, err := vb.call(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(workerErr).ToNot(BeEmpty()) + }) + + It("bounds the request by the context deadline", func() { + mc.scriptReply(vb.subject, vb.errorReply) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _, err := vb.call(ctx) + Expect(err).ToNot(HaveOccurred()) + t := lastTimeout() + Expect(t).To(BeNumerically(">", 0)) + Expect(t).To(BeNumerically("<=", 5*time.Second)) + }) + + It("falls back to the verb's default budget without a deadline", func() { + mc.scriptReply(vb.subject, vb.errorReply) + _, err := vb.call(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(lastTimeout()).To(Equal(vb.fallback)) + }) + + It("still issues the request when the context is cancelled but time is left", func() { + mc.scriptReply(vb.subject, vb.errorReply) + ctx, cancel := context.WithTimeout(context.Background(), time.Minute) + cancel() + _, err := vb.call(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(lastTimeout()).To(BeNumerically(">", 0)) + mc.mu.Lock() + defer mc.mu.Unlock() + Expect(mc.calls[len(mc.calls)-1].Subject).To(Equal(vb.subject)) + }) + }) + } +}) + +// failingSubscribeBus refuses every subscription, so a spec can check that a +// server passes the bus error up and names the verb that failed to start. +type failingSubscribeBus struct { + *testutil.FakeBus + err error +} + +func (f *failingSubscribeBus) Subscribe(string, func([]byte)) (messaging.Subscription, error) { + return nil, f.err +} + +func (f *failingSubscribeBus) QueueSubscribeReply(string, string, func([]byte, func([]byte))) (messaging.Subscription, error) { + return nil, f.err +} + +// The worker half must answer exactly as the agent worker always has: the +// frontend reads an undecodable request's reply text, the queue group decides +// which workers compete, and the handler context must outlive a shutdown. +var _ = Describe("NATS agent RPC server", func() { + var ( + bus *testutil.FakeBus + srv *NATSAgentRPCServer + ) + + BeforeEach(func() { + bus = testutil.NewFakeBus() + srv = NewNATSAgentRPCServer(bus, "n1") + }) + + type verb struct { + subject string + // serve registers a handler that records its context and decoded + // request, then returns reply (blocking on gate when it is non-nil). + serve func(reply any, gate <-chan struct{}, seen chan<- any, ctxs chan<- context.Context) error + // request is a decodable request and reply a response the handler + // returns; unmarshalErr decodes bad bytes the way the server must. + request any + reply any + unmarshalErr func(bad []byte) string + decodeReply func(data []byte) (string, error) + } + + verbs := []struct { + name string + v verb + }{ + {"mcp tool", verb{ + subject: messaging.SubjectMCPToolExecute, + serve: func(reply any, gate <-chan struct{}, seen chan<- any, ctxs chan<- context.Context) error { + return srv.ServeMCPTool(func(ctx context.Context, req mcpremote.MCPToolRequest) mcpremote.MCPToolResponse { + ctxs <- ctx + seen <- req + if gate != nil { + <-gate + } + return reply.(mcpremote.MCPToolResponse) + }) + }, + request: mcpremote.MCPToolRequest{ModelName: "m", ToolName: "t", Arguments: map[string]any{"a": "b"}}, + reply: mcpremote.MCPToolResponse{Result: "done", Error: "partial"}, + unmarshalErr: func(bad []byte) string { + var r mcpremote.MCPToolRequest + return json.Unmarshal(bad, &r).Error() + }, + decodeReply: func(data []byte) (string, error) { + var r mcpremote.MCPToolResponse + err := json.Unmarshal(data, &r) + return r.Error, err + }, + }}, + {"mcp discovery", verb{ + subject: messaging.SubjectMCPDiscovery, + serve: func(reply any, gate <-chan struct{}, seen chan<- any, ctxs chan<- context.Context) error { + return srv.ServeMCPDiscovery(func(ctx context.Context, req mcpremote.MCPDiscoveryRequest) mcpremote.MCPDiscoveryResponse { + ctxs <- ctx + seen <- req + if gate != nil { + <-gate + } + return reply.(mcpremote.MCPDiscoveryResponse) + }) + }, + request: mcpremote.MCPDiscoveryRequest{ModelName: "m"}, + reply: mcpremote.MCPDiscoveryResponse{ + Servers: []mcpremote.MCPServerInfo{{Name: "s", Type: "remote", Tools: []string{"t"}}}, + Tools: []mcpremote.MCPToolDef{{ServerName: "s", ToolName: "t"}}, + }, + unmarshalErr: func(bad []byte) string { + var r mcpremote.MCPDiscoveryRequest + return json.Unmarshal(bad, &r).Error() + }, + decodeReply: func(data []byte) (string, error) { + var r mcpremote.MCPDiscoveryResponse + err := json.Unmarshal(data, &r) + return r.Error, err + }, + }}, + } + + for _, entry := range verbs { + v := entry.v + Context(entry.name, func() { + var ( + seen chan any + ctxs chan context.Context + ) + + BeforeEach(func() { + seen = make(chan any, 4) + ctxs = make(chan context.Context, 4) + }) + + It("answers undecodable bytes with the unmarshal error and does not call the handler", func() { + Expect(v.serve(v.reply, nil, seen, ctxs)).To(Succeed()) + bad := []byte("{not json") + data, ok := bus.DeliverReply(v.subject, bad) + Expect(ok).To(BeTrue(), "a bad request must still be answered") + errText, err := v.decodeReply(data) + Expect(err).ToNot(HaveOccurred()) + Expect(errText).To(HavePrefix("unmarshal error: ")) + Expect(errText).To(Equal("unmarshal error: " + v.unmarshalErr(bad))) + Expect(seen).To(BeEmpty()) + }) + + It("calls the handler with the decoded request and sends its reply verbatim", func() { + Expect(v.serve(v.reply, nil, seen, ctxs)).To(Succeed()) + body, err := json.Marshal(v.request) + Expect(err).ToNot(HaveOccurred()) + data, ok := bus.DeliverReply(v.subject, body) + Expect(ok).To(BeTrue()) + Expect(seen).To(Receive(Equal(v.request))) + want, err := json.Marshal(v.reply) + Expect(err).ToNot(HaveOccurred()) + Expect(data).To(Equal(want)) + }) + + It("joins the agent-workers queue group", func() { + Expect(v.serve(v.reply, nil, seen, ctxs)).To(Succeed()) + Expect(bus.QueueGroups()).To(HaveKeyWithValue(v.subject, messaging.QueueAgentWorkers)) + Expect(messaging.QueueAgentWorkers).To(Equal("agent-workers")) + }) + + It("runs the handler on a context no parent can cancel", func() { + Expect(v.serve(v.reply, nil, seen, ctxs)).To(Succeed()) + body, err := json.Marshal(v.request) + Expect(err).ToNot(HaveOccurred()) + _, ok := bus.DeliverReply(v.subject, body) + Expect(ok).To(BeTrue()) + var ctx context.Context + Expect(ctxs).To(Receive(&ctx)) + Expect(ctx.Done()).To(BeNil(), "a cancellable context would abort an in-flight call on shutdown") + Expect(ctx.Err()).ToNot(HaveOccurred()) + _, hasDeadline := ctx.Deadline() + Expect(hasDeadline).To(BeFalse()) + }) + + It("handles deliveries one at a time on the delivery goroutine", func() { + gate := make(chan struct{}) + Expect(v.serve(v.reply, gate, seen, ctxs)).To(Succeed()) + body, err := json.Marshal(v.request) + Expect(err).ToNot(HaveOccurred()) + + // NATS delivers one subscription's messages in sequence on a + // single goroutine; this loop stands in for it. + var replies atomic.Int32 + done := make(chan struct{}) + go func() { + defer GinkgoRecover() + defer close(done) + for range 2 { + if _, ok := bus.DeliverReply(v.subject, body); ok { + replies.Add(1) + } + } + }() + + Eventually(seen).Should(Receive()) + Consistently(seen, 200*time.Millisecond).ShouldNot(Receive(), "the second request started before the first returned") + Expect(replies.Load()).To(BeZero(), "the reply must be sent by the delivery call itself") + + gate <- struct{}{} + Eventually(seen).Should(Receive()) + Expect(replies.Load()).To(Equal(int32(1))) + gate <- struct{}{} + Eventually(done).Should(BeClosed()) + Expect(replies.Load()).To(Equal(int32(2))) + }) + + It("returns a subscribe error naming the verb", func() { + boom := errors.New("permission denied") + failing := NewNATSAgentRPCServer(&failingSubscribeBus{FakeBus: bus, err: boom}, "n1") + var err error + if v.subject == messaging.SubjectMCPToolExecute { + err = failing.ServeMCPTool(func(context.Context, mcpremote.MCPToolRequest) mcpremote.MCPToolResponse { + return mcpremote.MCPToolResponse{} + }) + } else { + err = failing.ServeMCPDiscovery(func(context.Context, mcpremote.MCPDiscoveryRequest) mcpremote.MCPDiscoveryResponse { + return mcpremote.MCPDiscoveryResponse{} + }) + } + Expect(err).To(MatchError(boom)) + Expect(err.Error()).To(ContainSubstring(entry.name)) + }) + }) + } + + Context("backend stop", func() { + It("listens on the node's backend stop subject and never replies", func() { + var got []string + Expect(srv.ServeBackendStop(func(backend string) { got = append(got, backend) })).To(Succeed()) + + subject := messaging.SubjectNodeBackendStop("n1") + Expect(bus.Publish(subject, workerctl.BackendStopRequest{Backend: "llama"})).To(Succeed()) + Expect(got).To(Equal([]string{"llama"})) + + // A plain subscription has no reply to send: the frontend reads the + // resulting timeout as success, so a reply would change its outcome. + _, replied := bus.DeliverReply(subject, []byte(`{"backend":"llama"}`)) + Expect(replied).To(BeFalse()) + Expect(bus.QueueGroups()).ToNot(HaveKey(subject)) + }) + + It("ignores a body it cannot decode", func() { + called := false + Expect(srv.ServeBackendStop(func(string) { called = true })).To(Succeed()) + Expect(bus.Publish(messaging.SubjectNodeBackendStop("n1"), "not an object")).To(Succeed()) + Expect(called).To(BeFalse()) + }) + + It("does not hear another node's stop", func() { + called := false + Expect(srv.ServeBackendStop(func(string) { called = true })).To(Succeed()) + Expect(bus.Publish(messaging.SubjectNodeBackendStop("n2"), workerctl.BackendStopRequest{Backend: "llama"})).To(Succeed()) + Expect(called).To(BeFalse()) + }) + + It("returns a subscribe error naming the verb", func() { + boom := errors.New("permission denied") + failing := NewNATSAgentRPCServer(&failingSubscribeBus{FakeBus: bus, err: boom}, "n1") + err := failing.ServeBackendStop(func(string) {}) + Expect(err).To(MatchError(boom)) + Expect(err.Error()).To(ContainSubstring("backend stop")) + }) + }) +}) diff --git a/core/services/nodes/client_factory_test.go b/core/services/nodes/client_factory_test.go new file mode 100644 index 000000000..3c8dff7af --- /dev/null +++ b/core/services/nodes/client_factory_test.go @@ -0,0 +1,81 @@ +package nodes + +import ( + "context" + "sync" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + grpc "github.com/mudler/LocalAI/pkg/grpc" +) + +// recordingFactory records the node id, address and parallel flag of every +// client it builds, so a spec can assert that each consumer passes the node it +// is dialing and asks for the client it needs. +type recordingFactory struct { + mu sync.Mutex + seen []string + parallel []bool + next func() grpc.Backend +} + +func (f *recordingFactory) NewClient(nodeID, address string, parallel bool) grpc.Backend { + f.mu.Lock() + f.seen = append(f.seen, nodeID+"@"+address) + f.parallel = append(f.parallel, parallel) + f.mu.Unlock() + if f.next != nil { + return f.next() + } + return nil +} + +func (f *recordingFactory) calls() []string { + f.mu.Lock() + defer f.mu.Unlock() + return append([]string(nil), f.seen...) +} + +// parallelFlags returns the parallel argument of every build, in call order, +// so a spec can pin whether a consumer asks for a serialised client. +func (f *recordingFactory) parallelFlags() []bool { + f.mu.Lock() + defer f.mu.Unlock() + return append([]bool(nil), f.parallel...) +} + +var _ = Describe("Backend client construction carries the node id", func() { + It("hands the node id and the replica address to the factory from the health monitor", func() { + f := &recordingFactory{next: func() grpc.Backend { return &fakeBackendClient{healthy: true} }} + store := newFakeNodeHealthStore() + hm := newTestHealthMonitor(store, f, true, 30*time.Second) + hm.perModelHealthCheck = true + + store.addNode(makeTestNode("n1", "worker-1", "10.0.0.1:50051", StatusHealthy, freshTime())) + store.addNodeModel("n1", NodeModel{NodeID: "n1", ModelName: "m", Address: "10.0.0.1:50052"}) + + hm.doCheckAll(context.Background()) + + Expect(f.calls()).To(Equal([]string{"n1@10.0.0.1:50052"})) + }) + + It("hands the node id and the replica address to the factory from the router", func() { + f := &recordingFactory{next: func() grpc.Backend { return &stubBackend{healthResult: true} }} + reg := &fakeModelRouter{ + findAndLockNode: &BackendNode{ID: "n1", Name: "node-1", Address: "10.0.0.2:50051"}, + findAndLockNM: &NodeModel{NodeID: "n1", ModelName: "my-model", Address: "10.0.0.2:9001"}, + } + router := NewSmartRouter(reg, SmartRouterOptions{ClientFactory: f}) + + result, err := router.Route(context.Background(), "my-model", "models/my-model.gguf", "llama-cpp", "", nil, false) + Expect(err).ToNot(HaveOccurred()) + DeferCleanup(result.Release) + + // Route builds one client to probe the replica and one to serve it; + // both dial the same node, so every build must carry its id. + Expect(f.calls()).NotTo(BeEmpty()) + Expect(f.calls()).To(HaveEach("n1@10.0.0.2:9001")) + }) +}) diff --git a/core/services/nodes/control_errors_test.go b/core/services/nodes/control_errors_test.go new file mode 100644 index 000000000..7883e382e --- /dev/null +++ b/core/services/nodes/control_errors_test.go @@ -0,0 +1,63 @@ +package nodes + +import ( + "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" +) + +// The scheduler demotes a node on ErrNoRoute. Demoting on anything weaker +// condemns a node that is slow or busy, and a node that answers is present by +// demonstration, so a refusal must never look like a missing route. +var _ = Describe("Control request error classification", func() { + const nodeID = "11111111-2222-3333-4444-555555555555" + var ( + mc *scriptedMessagingClient + subject string + ) + + BeforeEach(func() { + mc = newScriptedMessagingClient() + subject = messaging.SubjectNodeBackendInstall(nodeID) + }) + + request := func() (*workerctl.BackendInstallReply, error) { + return controlRequestJSON[workerctl.BackendInstallRequest, workerctl.BackendInstallReply]( + mc, subject, workerctl.BackendInstallRequest{Backend: "b"}, time.Second) + } + + It("reports a subject nobody answers as ErrNoRoute", func() { + mc.scriptNoResponders(subject) + _, err := request() + Expect(errors.Is(err, ErrNoRoute)).To(BeTrue(), "got %v", err) + }) + + It("keeps the transport cause out of the unwrap chain", func() { + mc.scriptNoResponders(subject) + _, err := request() + Expect(errors.Is(err, nats.ErrNoResponders)).To(BeFalse(), + "consumers must match ErrNoRoute, never the carrier's own sentinel") + }) + + It("does not report a timeout as ErrNoRoute", func() { + mc.scriptErr(subject, nats.ErrTimeout) + _, err := request() + Expect(err).To(HaveOccurred()) + Expect(errors.Is(err, ErrNoRoute)).To(BeFalse()) + Expect(isNATSTimeout(err)).To(BeTrue()) + }) + + It("does not report a worker's own refusal as ErrNoRoute", func() { + mc.scriptReply(subject, workerctl.BackendInstallReply{Success: false, Error: "disk full"}) + reply, err := request() + Expect(err).ToNot(HaveOccurred()) + Expect(reply.Success).To(BeFalse()) + Expect(reply.Error).To(Equal("disk full")) + }) +}) diff --git a/core/services/nodes/control_nats.go b/core/services/nodes/control_nats.go new file mode 100644 index 000000000..eb6b41603 --- /dev/null +++ b/core/services/nodes/control_nats.go @@ -0,0 +1,59 @@ +package nodes + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/nats-io/nats.go" + + "github.com/mudler/LocalAI/core/services/messaging" +) + +// ErrNoRoute reports that a control command could not be delivered to a node: +// nothing is listening for it right now. +// +// It is a routing fact and says nothing about whether the node exists or is +// serving. The one reaction a consumer may have is the status-only demotion +// (MarkUnhealthy), which the next heartbeat reverses. Nothing may delete a +// node's model rows on it: a node that is registered and heartbeating can be +// unroutable for ordinary reasons, and reclaiming its models would evict healthy +// work. +// +// It is NOT returned for a timeout, for a transport fault, or for a worker that +// answered with a refusal. A node that answers is present by demonstration. +var ErrNoRoute = errors.New("nodes: no route to that node") + +// controlRequestJSON is messaging.RequestJSON with the carrier's failure mapped +// onto the conditions this package acts on. The carrier's own sentinel is kept +// out of the unwrap chain on purpose: it names an absence, and a consumer that +// matched on it would read absence as a fact about the node. +func controlRequestJSON[Req, Reply any](bus messaging.MessagingClient, subject string, req Req, timeout time.Duration) (*Reply, error) { + reply, err := messaging.RequestJSON[Req, Reply](bus, subject, req, timeout) + if err != nil && errors.Is(err, nats.ErrNoResponders) { + return nil, fmt.Errorf("%w: %v", ErrNoRoute, err) + } + return reply, err +} + +// isNATSTimeout returns true if err looks like a NATS request-reply timeout. +// nats.ErrTimeout is the canonical sentinel; context.DeadlineExceeded can +// also surface depending on the client's path; we accept both, plus a +// string-match fallback for clients that return a bare error. +func isNATSTimeout(err error) bool { + if errors.Is(err, nats.ErrTimeout) || errors.Is(err, context.DeadlineExceeded) { + return true + } + return err != nil && strings.Contains(err.Error(), "nats: timeout") +} + +// isStrictRequestTimeout matches only the carrier's own timeout sentinel. +// stopBackend reads a timeout as "an older worker performed the stop without +// replying" and reports success, so it must not inherit isNATSTimeout's wider +// net: a context deadline or a look-alike message there would turn a real +// failure into a silent success. +func isStrictRequestTimeout(err error) bool { + return errors.Is(err, nats.ErrTimeout) +} diff --git a/core/services/nodes/disk_headroom_test.go b/core/services/nodes/disk_headroom_test.go index f14abd829..17d35b550 100644 --- a/core/services/nodes/disk_headroom_test.go +++ b/core/services/nodes/disk_headroom_test.go @@ -9,8 +9,8 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "gorm.io/gorm" ) @@ -145,7 +145,7 @@ var _ = Describe("scheduling a model onto a cluster without disk headroom", func reg.findIdleNode = &BackendNode{ID: "n1", Name: "nvidia-thor", Address: "10.0.0.1:50051"} backend = &holdBackend{} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, } router = NewSmartRouter(reg, SmartRouterOptions{ Unloader: unloader, diff --git a/core/services/nodes/file_stager.go b/core/services/nodes/file_stager.go index 144c93165..71387426f 100644 --- a/core/services/nodes/file_stager.go +++ b/core/services/nodes/file_stager.go @@ -15,6 +15,10 @@ import ( // // 2. HTTPFileStager (fallback): Frontend pushes/pulls files directly over // HTTP to a small file transfer server on the backend node (no S3 needed). +// +// S3NATSFileStager returns ErrNoRoute when nothing is listening for the node; +// HTTPFileStager reports connection failures as ordinary errors. See ErrNoRoute +// for what a caller may do with it. type FileStager interface { // EnsureRemote ensures a local file is available on the remote node. // Returns the remote-local path. diff --git a/core/services/nodes/file_stager_dial_test.go b/core/services/nodes/file_stager_dial_test.go new file mode 100644 index 000000000..cb01eb856 --- /dev/null +++ b/core/services/nodes/file_stager_dial_test.go @@ -0,0 +1,101 @@ +package nodes + +import ( + "context" + "encoding/json" + "io" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("HTTPFileStager dialing", func() { + // recordingDialer sends every node to the listener named in routes and + // records which node each dial was for, so a spec can tell apart the node + // the stager asked for from the address it would have dialled itself. + recordingDialer := func(routes map[string]string) (WorkerNetDialerFor, func() []string) { + var mu sync.Mutex + var dials []string + dialFor := func(nodeID string) func(context.Context, string, string) (net.Conn, error) { + return func(ctx context.Context, network, _ string) (net.Conn, error) { + mu.Lock() + dials = append(dials, nodeID) + mu.Unlock() + return (&net.Dialer{}).DialContext(ctx, network, routes[nodeID]) + } + } + return dialFor, func() []string { + mu.Lock() + defer mu.Unlock() + return append([]string(nil), dials...) + } + } + + It("reaches the worker through the dialer of that node and reuses one client per node", func() { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodDelete || r.URL.Path != "/v1/files/ephemeral/request-id/audio/input.wav" { + http.Error(w, "unexpected request", http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusNoContent) + })) + DeferCleanup(srv.Close) + + dialFor, dials := recordingDialer(map[string]string{"n1": srv.Listener.Addr().String()}) + // httpAddrFor names a host that cannot resolve: only the dialer can reach it. + stager := NewHTTPFileStager(func(string) (string, error) { return "n1.worker.invalid:80", nil }, "tok", dialFor) + + key := "ephemeral/request-id/audio/input.wav" + Expect(stager.ReleaseRemote(context.Background(), "n1", key)).To(Succeed()) + Expect(stager.ReleaseRemote(context.Background(), "n1", key)).To(Succeed()) + + Expect(dials()).To(Equal([]string{"n1"}), "two calls to one node must reuse one client and one connection") + }) + + It("keeps two nodes that report the same address on their own dialers", func() { + // Each fake worker answers HEAD with 404 (no copy yet) and a PUT with + // the path it stored, prefixed with its own name, so the returned path + // shows which worker actually received the bytes. + worker := func(name string) *httptest.Server { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodHead: + w.WriteHeader(http.StatusNotFound) + case http.MethodPut: + _, _ = io.Copy(io.Discard, r.Body) + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]string{"local_path": "/" + name + r.URL.Path}) + default: + http.Error(w, "unexpected request", http.StatusBadRequest) + } + })) + DeferCleanup(srv.Close) + return srv + } + srvA, srvB := worker("a"), worker("b") + + dialFor, dials := recordingDialer(map[string]string{ + "n1": srvA.Listener.Addr().String(), + "n2": srvB.Listener.Addr().String(), + }) + stager := NewHTTPFileStager(func(string) (string, error) { return "shared.worker.invalid:80", nil }, "", dialFor) + + localPath := filepath.Join(GinkgoT().TempDir(), "model.bin") + Expect(os.WriteFile(localPath, []byte("weights"), 0o600)).To(Succeed()) + + pathA, err := stager.EnsureRemote(context.Background(), "n1", localPath, "models/model.bin") + Expect(err).NotTo(HaveOccurred()) + pathB, err := stager.EnsureRemote(context.Background(), "n2", localPath, "models/model.bin") + Expect(err).NotTo(HaveOccurred()) + + Expect(pathA).To(Equal("/a/v1/files/models/model.bin")) + Expect(pathB).To(Equal("/b/v1/files/models/model.bin")) + Expect(dials()).To(ConsistOf("n1", "n2"), "the probe, resume HEAD and PUT of one call share that node's connection") + }) +}) diff --git a/core/services/nodes/file_stager_http.go b/core/services/nodes/file_stager_http.go index 0a64db5ab..4cfde4bfd 100644 --- a/core/services/nodes/file_stager_http.go +++ b/core/services/nodes/file_stager_http.go @@ -32,7 +32,9 @@ import ( type HTTPFileStager struct { httpAddrFor func(nodeID string) (string, error) token string - client *http.Client + dialFor WorkerNetDialerFor + clientsMu sync.Mutex + clients map[string]*http.Client responseTimeout time.Duration // timeout waiting for server response after upload maxRetries int // number of retry attempts for transient failures } @@ -40,7 +42,8 @@ type HTTPFileStager struct { // NewHTTPFileStager creates a new HTTP file stager. // httpAddrFor should return the HTTP address (host:port) for the given node ID. // token is the registration token used for authentication. -func NewHTTPFileStager(httpAddrFor func(nodeID string) (string, error), token string) *HTTPFileStager { +// dialFor returns the dial function that reaches a given node's HTTP server. +func NewHTTPFileStager(httpAddrFor func(nodeID string) (string, error), token string, dialFor WorkerNetDialerFor) *HTTPFileStager { responseTimeout := 30 * time.Minute if v := os.Getenv("LOCALAI_FILE_TRANSFER_TIMEOUT"); v != "" { if d, err := time.ParseDuration(v); err == nil { @@ -55,11 +58,28 @@ func NewHTTPFileStager(httpAddrFor func(nodeID string) (string, error), token st } } + return &HTTPFileStager{ + httpAddrFor: httpAddrFor, + token: token, + dialFor: dialFor, + clients: map[string]*http.Client{}, + responseTimeout: responseTimeout, + maxRetries: maxRetries, + } +} + +// clientFor returns the HTTP client that reaches nodeID. Clients are per node +// rather than shared because the idle pool is keyed by host:port only: two +// workers reporting the same address (NAT, loopback) would otherwise be handed +// each other's connections once the dialer routes by node. +func (h *HTTPFileStager) clientFor(nodeID string) *http.Client { + h.clientsMu.Lock() + defer h.clientsMu.Unlock() + if c, ok := h.clients[nodeID]; ok { + return c + } transport := &http.Transport{ - DialContext: (&net.Dialer{ - Timeout: 30 * time.Second, - KeepAlive: 15 * time.Second, // aggressive keepalive for LAN transfers - }).DialContext, + DialContext: h.dialFor(nodeID), ForceAttemptHTTP2: false, // HTTP/2 flow control can stall large uploads MaxIdleConns: 10, IdleConnTimeout: 90 * time.Second, @@ -68,19 +88,14 @@ func NewHTTPFileStager(httpAddrFor func(nodeID string) (string, error), token st WriteBufferSize: 256 << 10, // 256 KB ReadBufferSize: 256 << 10, // 256 KB } - - return &HTTPFileStager{ - httpAddrFor: httpAddrFor, - token: token, - // No Timeout set — for large uploads, http.Client.Timeout covers the - // entire request lifecycle including the body upload. If it fires - // mid-write, Go closes the connection causing "connection reset by peer" - // on the server. Instead we use ResponseHeaderTimeout on the transport - // to cover only the wait-for-server-response phase. - client: httpclient.New(httpclient.WithTransport(transport)), - responseTimeout: responseTimeout, - maxRetries: maxRetries, - } + // No Timeout set: for large uploads, http.Client.Timeout covers the + // entire request lifecycle including the body upload. If it fires + // mid-write, Go closes the connection causing "connection reset by peer" + // on the server. Instead we use ResponseHeaderTimeout on the transport + // to cover only the wait-for-server-response phase. + c := httpclient.New(httpclient.WithTransport(transport)) + h.clients[nodeID] = c + return c } // ReleaseRemote removes one exact ephemeral key from a backend node. @@ -100,7 +115,7 @@ func (h *HTTPFileStager) ReleaseRemote(ctx context.Context, nodeID, key string) if h.token != "" { req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := h.clientFor(nodeID).Do(req) if err != nil { return fmt.Errorf("releasing %q from node %s: %w", key, nodeID, err) } @@ -138,7 +153,7 @@ func (h *HTTPFileStager) ReleaseRemoteRequest(ctx context.Context, nodeID, reque if h.token != "" { req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := h.clientFor(nodeID).Do(req) if err != nil { return fmt.Errorf("releasing request inputs from node %s: %w", nodeID, err) } @@ -170,9 +185,11 @@ func (h *HTTPFileStager) EnsureRemote(ctx context.Context, nodeID, localPath, ke if err != nil { return "", fmt.Errorf("resolving HTTP address for node %s: %w", nodeID, err) } + // Fetched once per call so every retry reuses the same connection pool. + client := h.clientFor(nodeID) // Probe: check if the remote already has the file with matching content hash. - if remotePath, ok, probeErr := h.probeExisting(ctx, addr, localPath, key); probeErr != nil { + if remotePath, ok, probeErr := h.probeExisting(ctx, client, addr, localPath, key); probeErr != nil { return "", fmt.Errorf("claiming existing file on node %s: %w", nodeID, probeErr) } else if ok { xlog.Info("Upload skipped (file already exists with matching hash)", "node", nodeID, "key", key, "remotePath", remotePath) @@ -232,9 +249,9 @@ func (h *HTTPFileStager) EnsureRemote(ctx context.Context, nodeID, localPath, ke // matching ours unlocks resume from the reported size; any other // outcome (missing file, hash mismatch, partial-of-different-file) // resets to 0 and uploads the entire file. - startOffset := h.resumeOffset(resumeCtx, addr, key, localHash, fileSize) + startOffset := h.resumeOffset(resumeCtx, client, addr, key, localHash, fileSize) - result, err := h.doUpload(ctx, resumeCtx, addr, nodeID, localPath, key, url, fileSize, startOffset, localHash) + result, err := h.doUpload(ctx, resumeCtx, client, addr, nodeID, localPath, key, url, fileSize, startOffset, localHash) if err == nil { if attempt > 1 { xlog.Info("File upload succeeded after retry", "node", nodeID, "file", filepath.Base(localPath), "attempt", attempt) @@ -321,7 +338,7 @@ func nextBackoff(attempt int) time.Duration { // different target hash). It returns the server-reported size when the // server's X-Target-SHA256 matches our expected final hash AND the size is // strictly less than the local file size. -func (h *HTTPFileStager) resumeOffset(ctx context.Context, addr, key, localHash string, fileSize int64) int64 { +func (h *HTTPFileStager) resumeOffset(ctx context.Context, client *http.Client, addr, key, localHash string, fileSize int64) int64 { if localHash == "" || fileSize <= 0 { return 0 } @@ -333,7 +350,7 @@ func (h *HTTPFileStager) resumeOffset(ctx context.Context, addr, key, localHash if h.token != "" { req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := client.Do(req) if err != nil { return 0 } @@ -366,7 +383,7 @@ func (h *HTTPFileStager) resumeOffset(ctx context.Context, addr, key, localHash // the bytes from startOffset to fileSize-1. The outerCtx is the long-lived // resume budget; reqCtx is what's bound to the request (currently the same as // the parent ctx, since http.Client doesn't expose a per-request timeout). -func (h *HTTPFileStager) doUpload(ctx, outerCtx context.Context, addr, nodeID, localPath, key, url string, fileSize, startOffset int64, expectedHash string) (string, error) { +func (h *HTTPFileStager) doUpload(ctx, outerCtx context.Context, client *http.Client, addr, nodeID, localPath, key, url string, fileSize, startOffset int64, expectedHash string) (string, error) { if startOffset < 0 || startOffset > fileSize { startOffset = 0 } @@ -421,7 +438,7 @@ func (h *HTTPFileStager) doUpload(ctx, outerCtx context.Context, addr, nodeID, l req.Header.Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", startOffset, fileSize-1, fileSize)) } - resp, err := h.client.Do(req) + resp, err := client.Do(req) if err != nil { xlog.Error("File upload failed", "node", nodeID, "file", filepath.Base(localPath), "size", humanFileSize(fileSize), "offset", startOffset, "error", err) @@ -526,7 +543,7 @@ func isTransientError(err error) bool { // upload can be skipped. HEAD and hash errors fall through to a normal PUT. // Matching ephemeral files are claimed first; a 404 or 405 claim response // identifies an older worker and also falls through to PUT. -func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key string) (string, bool, error) { +func (h *HTTPFileStager) probeExisting(ctx context.Context, client *http.Client, addr, localPath, key string) (string, bool, error) { url := fmt.Sprintf("http://%s/v1/files/%s", addr, key) req, err := http.NewRequestWithContext(ctx, http.MethodHead, url, nil) @@ -537,7 +554,7 @@ func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := client.Do(req) if err != nil { return "", false, nil } @@ -568,7 +585,7 @@ func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key } if strings.HasPrefix(key, "ephemeral/") { - claimed, err := h.claimExisting(ctx, addr, key) + claimed, err := h.claimExisting(ctx, client, addr, key) if err != nil { return "", false, err } @@ -580,7 +597,7 @@ func (h *HTTPFileStager) probeExisting(ctx context.Context, addr, localPath, key return remotePath, true, nil } -func (h *HTTPFileStager) claimExisting(ctx context.Context, addr, key string) (bool, error) { +func (h *HTTPFileStager) claimExisting(ctx context.Context, client *http.Client, addr, key string) (bool, error) { claimURL := (&url.URL{ Scheme: "http", Host: addr, @@ -594,7 +611,7 @@ func (h *HTTPFileStager) claimExisting(ctx context.Context, addr, key string) (b if h.token != "" { req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := client.Do(req) if err != nil { return false, fmt.Errorf("claiming %q: %w", key, err) } @@ -804,7 +821,7 @@ func (h *HTTPFileStager) FetchRemoteByKey(ctx context.Context, nodeID, key, loca req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := h.clientFor(nodeID).Do(req) if err != nil { return fmt.Errorf("downloading from node %s: %w", nodeID, err) } @@ -860,7 +877,7 @@ func (h *HTTPFileStager) AllocRemoteTemp(ctx context.Context, nodeID string) (st req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := h.clientFor(nodeID).Do(req) if err != nil { return "", fmt.Errorf("allocating temp file on node %s: %w", nodeID, err) } @@ -901,7 +918,7 @@ func (h *HTTPFileStager) ListRemoteDir(ctx context.Context, nodeID, keyPrefix st req.Header.Set("Authorization", "Bearer "+h.token) } - resp, err := h.client.Do(req) + resp, err := h.clientFor(nodeID).Do(req) if err != nil { return nil, fmt.Errorf("listing dir on node %s: %w", nodeID, err) } diff --git a/core/services/nodes/file_stager_release_test.go b/core/services/nodes/file_stager_release_test.go index 519ef6287..b3932e356 100644 --- a/core/services/nodes/file_stager_release_test.go +++ b/core/services/nodes/file_stager_release_test.go @@ -14,6 +14,7 @@ import ( "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/storage" + "github.com/mudler/LocalAI/core/services/workerctl" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -77,7 +78,7 @@ var _ = Describe("File stager exact-key release", func() { Expect(err).NotTo(HaveOccurred()) return NewHTTPFileStager(func(string) (string, error) { return listener.Addr().String(), nil - }, token), func() { + }, token, DirectWorkerNetDialer()), func() { Expect(server.Shutdown(context.Background())).To(Succeed()) } } @@ -161,7 +162,7 @@ var _ = Describe("File stager exact-key release", func() { DeferCleanup(server.Close) stager := NewHTTPFileStager(func(string) (string, error) { return strings.TrimPrefix(server.URL, "http://"), nil - }, "") + }, "", DirectWorkerNetDialer()) keys := []string{ "ephemeral/audio/request-id/input.wav", "ephemeral/images/request-id/frame.jpg", @@ -178,7 +179,7 @@ var _ = Describe("File stager exact-key release", func() { stager := NewHTTPFileStager(func(string) (string, error) { resolved = true return "127.0.0.1:1", nil - }, "token") + }, "token", DirectWorkerNetDialer()) for _, key := range []string{ "models/model.gguf", @@ -247,7 +248,7 @@ var _ = Describe("File stager exact-key release", func() { Expect(client.requestCalled).To(BeTrue()) Expect(client.subject).To(Equal(messaging.SubjectNodeFilesRelease("node.one"))) - var payload fileReleaseRequest + var payload workerctl.FileReleaseRequest Expect(json.Unmarshal(client.payload, &payload)).To(Succeed()) Expect(payload.Key).To(Equal(key)) exists, err := store.Exists(context.Background(), key) @@ -274,7 +275,7 @@ var _ = Describe("File stager exact-key release", func() { Expect(stager.ReleaseRemoteRequest(context.Background(), "node.one", "request-id", keys)).To(Succeed()) Expect(client.requestCount).To(Equal(1)) - var payload fileReleaseRequest + var payload workerctl.FileReleaseRequest Expect(json.Unmarshal(client.payload, &payload)).To(Succeed()) Expect(payload.Key).To(BeEmpty()) Expect(payload.RequestID).To(Equal("request-id")) @@ -301,7 +302,7 @@ var _ = Describe("File stager exact-key release", func() { Expect(client.requestCount).To(Equal(1)) Expect(len(client.payload)).To(BeNumerically("<", 128)) - var payload fileReleaseRequest + var payload workerctl.FileReleaseRequest Expect(json.Unmarshal(client.payload, &payload)).To(Succeed()) Expect(payload.RequestID).To(Equal("request-id")) }) diff --git a/core/services/nodes/file_stager_s3.go b/core/services/nodes/file_stager_s3.go index fff6d1859..5f1c2bd55 100644 --- a/core/services/nodes/file_stager_s3.go +++ b/core/services/nodes/file_stager_s3.go @@ -8,6 +8,7 @@ import ( "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/storage" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/xlog" ) @@ -28,52 +29,6 @@ func NewS3NATSFileStager(fm *storage.FileManager, nats messaging.MessagingClient return &S3NATSFileStager{fm: fm, nats: nats} } -// NATS request/reply message types - -type fileEnsureRequest struct { - Key string `json:"key"` -} - -type fileEnsureReply struct { - LocalPath string `json:"local_path"` - Error string `json:"error,omitempty"` -} - -type fileStageRequest struct { - LocalPath string `json:"local_path"` - Key string `json:"key"` -} - -type fileStageReply struct { - Key string `json:"key"` - Error string `json:"error,omitempty"` -} - -type fileReleaseRequest struct { - Key string `json:"key,omitempty"` - RequestID string `json:"request_id,omitempty"` -} - -type fileReleaseReply struct { - Error string `json:"error,omitempty"` -} - -type fileTempRequest struct{} - -type fileTempReply struct { - LocalPath string `json:"local_path"` - Error string `json:"error,omitempty"` -} - -type fileListDirRequest struct { - KeyPrefix string `json:"key_prefix"` -} - -type fileListDirReply struct { - Files []string `json:"files"` - Error string `json:"error,omitempty"` -} - // EnsureRemote uploads a local file to S3 (if not already there) and sends // a NATS request-reply to the backend node to download it locally. func (s *S3NATSFileStager) EnsureRemote(ctx context.Context, nodeID, localPath, key string) (string, error) { @@ -94,7 +49,7 @@ func (s *S3NATSFileStager) EnsureRemote(ctx context.Context, nodeID, localPath, // Send NATS request-reply to backend subject := messaging.SubjectNodeFilesEnsure(nodeID) - reply, err := messaging.RequestJSON[fileEnsureRequest, fileEnsureReply](s.nats, subject, fileEnsureRequest{Key: key}, 10*time.Minute) + reply, err := controlRequestJSON[workerctl.FileEnsureRequest, workerctl.FileEnsureReply](s.nats, subject, workerctl.FileEnsureRequest{Key: key}, 10*time.Minute) if err != nil { return "", err } @@ -124,7 +79,7 @@ func (s *S3NATSFileStager) FetchRemoteByKey(ctx context.Context, nodeID, key, lo func (s *S3NATSFileStager) fetchRemoteWithKey(ctx context.Context, nodeID, remotePath, key, localDst string, cleanup bool) error { subject := messaging.SubjectNodeFilesStage(nodeID) - reply, err := messaging.RequestJSON[fileStageRequest, fileStageReply](s.nats, subject, fileStageRequest{LocalPath: remotePath, Key: key}, 10*time.Minute) + reply, err := controlRequestJSON[workerctl.FileStageRequest, workerctl.FileStageReply](s.nats, subject, workerctl.FileStageRequest{LocalPath: remotePath, Key: key}, 10*time.Minute) if err != nil { return err } @@ -154,7 +109,7 @@ func (s *S3NATSFileStager) fetchRemoteWithKey(ctx context.Context, nodeID, remot // AllocRemoteTemp asks the backend to allocate a temp file via NATS request-reply. func (s *S3NATSFileStager) AllocRemoteTemp(ctx context.Context, nodeID string) (string, error) { subject := messaging.SubjectNodeFilesTemp(nodeID) - reply, err := messaging.RequestJSON[fileTempRequest, fileTempReply](s.nats, subject, fileTempRequest{}, 30*time.Second) + reply, err := controlRequestJSON[workerctl.FileTempRequest, workerctl.FileTempReply](s.nats, subject, workerctl.FileTempRequest{}, 30*time.Second) if err != nil { return "", err } @@ -167,7 +122,7 @@ func (s *S3NATSFileStager) AllocRemoteTemp(ctx context.Context, nodeID string) ( func (s *S3NATSFileStager) ListRemoteDir(ctx context.Context, nodeID, keyPrefix string) ([]string, error) { subject := messaging.SubjectNodeFilesListDir(nodeID) - reply, err := messaging.RequestJSON[fileListDirRequest, fileListDirReply](s.nats, subject, fileListDirRequest{KeyPrefix: keyPrefix}, 30*time.Second) + reply, err := controlRequestJSON[workerctl.FileListDirRequest, workerctl.FileListDirReply](s.nats, subject, workerctl.FileListDirRequest{KeyPrefix: keyPrefix}, 30*time.Second) if err != nil { return nil, err } @@ -181,7 +136,7 @@ func (s *S3NATSFileStager) ListRemoteDir(ctx context.Context, nodeID, keyPrefix // StageRemoteToStore tells the backend to upload a local file to S3. func (s *S3NATSFileStager) StageRemoteToStore(ctx context.Context, nodeID, remotePath, key string) error { subject := messaging.SubjectNodeFilesStage(nodeID) - reply, err := messaging.RequestJSON[fileStageRequest, fileStageReply](s.nats, subject, fileStageRequest{LocalPath: remotePath, Key: key}, 10*time.Minute) + reply, err := controlRequestJSON[workerctl.FileStageRequest, workerctl.FileStageReply](s.nats, subject, workerctl.FileStageRequest{LocalPath: remotePath, Key: key}, 10*time.Minute) if err != nil { return err } @@ -198,7 +153,7 @@ func (s *S3NATSFileStager) ReleaseRemote(ctx context.Context, nodeID, key string if err := validateEphemeralReleaseKey(key); err != nil { return err } - if err := s.releaseWorkerKeys(ctx, nodeID, fileReleaseRequest{Key: key}); err != nil { + if err := s.releaseWorkerKeys(ctx, nodeID, workerctl.FileReleaseRequest{Key: key}); err != nil { return err } if err := s.fm.Delete(ctx, key); err != nil { @@ -214,7 +169,7 @@ func (s *S3NATSFileStager) ReleaseRemoteRequest(ctx context.Context, nodeID, req if err := validateEphemeralRequestRelease(requestID, keys); err != nil { return err } - if err := s.releaseWorkerKeys(ctx, nodeID, fileReleaseRequest{RequestID: requestID}); err != nil { + if err := s.releaseWorkerKeys(ctx, nodeID, workerctl.FileReleaseRequest{RequestID: requestID}); err != nil { var fallbackErrors []error for _, key := range keys { if fallbackErr := s.ReleaseRemote(ctx, nodeID, key); fallbackErr != nil { @@ -235,7 +190,7 @@ func (s *S3NATSFileStager) ReleaseRemoteRequest(ctx context.Context, nodeID, req return errors.Join(deleteErrors...) } -func (s *S3NATSFileStager) releaseWorkerKeys(ctx context.Context, nodeID string, request fileReleaseRequest) error { +func (s *S3NATSFileStager) releaseWorkerKeys(ctx context.Context, nodeID string, request workerctl.FileReleaseRequest) error { if err := ctx.Err(); err != nil { return err } @@ -247,7 +202,7 @@ func (s *S3NATSFileStager) releaseWorkerKeys(ctx context.Context, nodeID string, } timeout = min(timeout, remaining) } - reply, err := messaging.RequestJSON[fileReleaseRequest, fileReleaseReply]( + reply, err := controlRequestJSON[workerctl.FileReleaseRequest, workerctl.FileReleaseReply]( s.nats, messaging.SubjectNodeFilesRelease(nodeID), request, diff --git a/core/services/nodes/file_stager_verify_deadline_test.go b/core/services/nodes/file_stager_verify_deadline_test.go index 0827bbecd..45b979ac1 100644 --- a/core/services/nodes/file_stager_verify_deadline_test.go +++ b/core/services/nodes/file_stager_verify_deadline_test.go @@ -65,7 +65,7 @@ var _ = Describe("staging verify phase and the cold-load stall window", func() { return "", err } return u.Host, nil - }, "") + }, "", DirectWorkerNetDialer()) } It("survives a run of verified-and-skipped shards that upload no bytes at all", func() { diff --git a/core/services/nodes/file_staging_sound_detection_test.go b/core/services/nodes/file_staging_sound_detection_test.go index 2fc6fb7e3..865dfe18d 100644 --- a/core/services/nodes/file_staging_sound_detection_test.go +++ b/core/services/nodes/file_staging_sound_detection_test.go @@ -34,7 +34,7 @@ func (s *soundStagingFailure) ReleaseRemote(context.Context, string, string) err type soundRouteFactory struct{ client grpc.Backend } -func (f *soundRouteFactory) NewClient(string, bool) grpc.Backend { return f.client } +func (f *soundRouteFactory) NewClient(string, string, bool) grpc.Backend { return f.client } var _ = Describe("FileStagingClient sound detection", func() { It("stages sound audio through the client returned by SmartRouter.Route", func(ctx SpecContext) { diff --git a/core/services/nodes/file_transfer_finalize_test.go b/core/services/nodes/file_transfer_finalize_test.go index bc0753b0e..635210357 100644 --- a/core/services/nodes/file_transfer_finalize_test.go +++ b/core/services/nodes/file_transfer_finalize_test.go @@ -41,7 +41,7 @@ var _ = Describe("Recovering unfinished file finalization", func() { DeferCleanup(server.Close) stager := NewHTTPFileStager(func(string) (string, error) { return strings.TrimPrefix(server.URL, "http://"), nil - }, "") + }, "", DirectWorkerNetDialer()) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() result, err := stager.EnsureRemote(ctx, "worker", local, "model.bin") diff --git a/core/services/nodes/file_transfer_server_test.go b/core/services/nodes/file_transfer_server_test.go index 918379c5c..ee8486833 100644 --- a/core/services/nodes/file_transfer_server_test.go +++ b/core/services/nodes/file_transfer_server_test.go @@ -645,7 +645,7 @@ var _ = Describe("FileTransferServer", func() { DeferCleanup(ts.Close) stager := NewHTTPFileStager(func(string) (string, error) { return strings.TrimPrefix(ts.URL, "http://"), nil - }, "") + }, "", DirectWorkerNetDialer()) for range 2 { path, err := stager.EnsureRemote(context.Background(), "node-1", localPath, key) @@ -675,7 +675,7 @@ var _ = Describe("FileTransferServer", func() { DeferCleanup(ts.Close) stager := NewHTTPFileStager(func(string) (string, error) { return strings.TrimPrefix(ts.URL, "http://"), nil - }, "") + }, "", DirectWorkerNetDialer()) path, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "ephemeral/audio/request/input.wav") @@ -714,7 +714,7 @@ var _ = Describe("FileTransferServer", func() { DeferCleanup(ts.Close) stager := NewHTTPFileStager(func(string) (string, error) { return strings.TrimPrefix(ts.URL, "http://"), nil - }, "") + }, "", DirectWorkerNetDialer()) backend := &lifecycleBackend{} client := NewFileStagingClient(backend, stager, "node-1") @@ -750,7 +750,7 @@ var _ = Describe("FileTransferServer", func() { DeferCleanup(ts.Close) stager := NewHTTPFileStager(func(string) (string, error) { return strings.TrimPrefix(ts.URL, "http://"), nil - }, "") + }, "", DirectWorkerNetDialer()) path, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "models/tracking/model.bin") @@ -780,7 +780,7 @@ var _ = Describe("FileTransferServer", func() { addr := strings.TrimPrefix(ts.URL, "http://") stager := NewHTTPFileStager(func(nodeID string) (string, error) { return addr, nil - }, "tok") + }, "tok", DirectWorkerNetDialer()) remotePath, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "present.bin") Expect(err).ToNot(HaveOccurred()) @@ -809,7 +809,7 @@ var _ = Describe("FileTransferServer", func() { addr := strings.TrimPrefix(ts.URL, "http://") stager := NewHTTPFileStager(func(nodeID string) (string, error) { return addr, nil - }, "tok") + }, "tok", DirectWorkerNetDialer()) remotePath, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "changed.bin") Expect(err).ToNot(HaveOccurred()) @@ -838,7 +838,7 @@ var _ = Describe("FileTransferServer", func() { addr := strings.TrimPrefix(ts.URL, "http://") stager := NewHTTPFileStager(func(nodeID string) (string, error) { return addr, nil - }, "tok") + }, "tok", DirectWorkerNetDialer()) remotePath, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "new.bin") Expect(err).ToNot(HaveOccurred()) @@ -874,7 +874,7 @@ var _ = Describe("FileTransferServer", func() { addr := strings.TrimPrefix(ts.URL, "http://") stager := NewHTTPFileStager(func(nodeID string) (string, error) { return addr, nil - }, "") + }, "", DirectWorkerNetDialer()) remotePath, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "compat.bin") Expect(err).ToNot(HaveOccurred()) @@ -1091,7 +1091,7 @@ var _ = Describe("FileTransferServer", func() { addr := strings.TrimPrefix(ts.URL, "http://") stager := NewHTTPFileStager(func(nodeID string) (string, error) { return addr, nil - }, "tok") + }, "tok", DirectWorkerNetDialer()) remotePath, err := stager.EnsureRemote(context.Background(), "node-1", localPath, "resume.bin") Expect(err).ToNot(HaveOccurred()) @@ -1189,7 +1189,7 @@ var _ = Describe("FileTransferServer", func() { addr := strings.TrimPrefix(ts.URL, "http://") stager := NewHTTPFileStager(func(nodeID string) (string, error) { return addr, nil - }, "tok") + }, "tok", DirectWorkerNetDialer()) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() diff --git a/core/services/nodes/health.go b/core/services/nodes/health.go index b82e57f91..970dfda0d 100644 --- a/core/services/nodes/health.go +++ b/core/services/nodes/health.go @@ -189,7 +189,7 @@ func (hm *HealthMonitor) doCheckAll(ctx context.Context) { if m.Address == "" || m.Address == node.Address { continue } - mClient := hm.clientFactory.NewClient(m.Address, false) + mClient := hm.clientFactory.NewClient(node.ID, m.Address, false) mCheckCtx, mCancel := context.WithTimeout(ctx, 5*time.Second) ok, _ := mClient.HealthCheck(mCheckCtx) mCancel() diff --git a/core/services/nodes/health_mock_test.go b/core/services/nodes/health_mock_test.go index d8e74004a..fc00e6bb1 100644 --- a/core/services/nodes/health_mock_test.go +++ b/core/services/nodes/health_mock_test.go @@ -319,7 +319,7 @@ func (f *fakeBackendClientFactory) setClient(address string, c *fakeBackendClien f.clients[address] = c } -func (f *fakeBackendClientFactory) NewClient(address string, _ bool) grpc.Backend { +func (f *fakeBackendClientFactory) NewClient(_, address string, _ bool) grpc.Backend { f.mu.Lock() defer f.mu.Unlock() if c, ok := f.clients[address]; ok { diff --git a/core/services/nodes/install_progress_publisher.go b/core/services/nodes/install_progress_publisher.go index 60eacb711..001313ac4 100644 --- a/core/services/nodes/install_progress_publisher.go +++ b/core/services/nodes/install_progress_publisher.go @@ -4,43 +4,42 @@ import ( "sync" "time" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" ) -// DebouncedInstallProgressPublisher buffers backend-install download ticks -// and publishes them to the per-op NATS progress subject at most once per -// `interval`. Always publishes the final event on Flush so the UI sees the -// terminal percentage. +// DebouncedInstallProgressSink buffers backend-install download ticks and +// hands them to emit at most once per `interval`. Always emits the final +// event on Flush so the UI sees the terminal percentage. The debounce lives +// here rather than in the carrier behind emit, so every carrier sees the same +// bounded event rate. // // Behavior: leading-edge debounce. The first OnDownload after a quiet window -// publishes immediately; subsequent ticks within `interval` only buffer the +// emits immediately; subsequent ticks within `interval` only buffer the // latest event, which is then emitted via a single trailing timer. This // keeps the wire chatter bounded (~4 events per second at 250ms) while // still surfacing every meaningful percentage jump. // -// Lock ordering: never hold p.mu across a Publish call. Publish hits the -// NATS client which may block on a slow link, and we don't want a stalled -// network to stall the underlying gallery download loop. -type DebouncedInstallProgressPublisher struct { - mu sync.Mutex - client messaging.MessagingClient - subject string - nodeID string - opID string - backend string - interval time.Duration - lastPublishedAt time.Time - pending *messaging.BackendInstallProgressEvent - timer *time.Timer +// Lock ordering: never hold p.mu across an emit call. emit may block on a +// slow link, and we don't want a stalled network to stall the underlying +// gallery download loop. +type DebouncedInstallProgressSink struct { + mu sync.Mutex + emit func(workerctl.BackendInstallProgressEvent) + nodeID string + opID string + backend string + interval time.Duration + lastEmittedAt time.Time + pending *workerctl.BackendInstallProgressEvent + timer *time.Timer } -// NewDebouncedInstallProgressPublisher constructs a publisher for one -// install operation. interval is the leading-edge debounce window -// (~250ms in production). -func NewDebouncedInstallProgressPublisher(client messaging.MessagingClient, nodeID, opID, backend string, interval time.Duration) *DebouncedInstallProgressPublisher { - return &DebouncedInstallProgressPublisher{ - client: client, - subject: messaging.SubjectNodeBackendInstallProgress(nodeID, opID), +// NewDebouncedInstallProgressSink constructs a sink for one install +// operation. interval is the leading-edge debounce window (~250ms in +// production). +func NewDebouncedInstallProgressSink(emit func(workerctl.BackendInstallProgressEvent), nodeID, opID, backend string, interval time.Duration) *DebouncedInstallProgressSink { + return &DebouncedInstallProgressSink{ + emit: emit, nodeID: nodeID, opID: opID, backend: backend, @@ -51,8 +50,8 @@ func NewDebouncedInstallProgressPublisher(client messaging.MessagingClient, node // OnDownload is the callback shape gallery.InstallBackendFromGallery and // galleryop.InstallExternalBackend pass into the worker. Each invocation // represents a single tick from the underlying io.Reader copy loop. -func (p *DebouncedInstallProgressPublisher) OnDownload(file, current, total string, percentage float64) { - ev := messaging.BackendInstallProgressEvent{ +func (p *DebouncedInstallProgressSink) OnDownload(file, current, total string, percentage float64) { + ev := workerctl.BackendInstallProgressEvent{ OpID: p.opID, NodeID: p.nodeID, Backend: p.backend, @@ -60,52 +59,52 @@ func (p *DebouncedInstallProgressPublisher) OnDownload(file, current, total stri Current: current, Total: total, Percentage: percentage, - Phase: messaging.PhaseDownloading, + Phase: workerctl.PhaseDownloading, } p.mu.Lock() now := time.Now() - if p.lastPublishedAt.IsZero() || now.Sub(p.lastPublishedAt) >= p.interval { - // Leading edge: publish immediately. - p.lastPublishedAt = now + if p.lastEmittedAt.IsZero() || now.Sub(p.lastEmittedAt) >= p.interval { + // Leading edge: emit immediately. + p.lastEmittedAt = now p.pending = nil p.mu.Unlock() - _ = p.client.Publish(p.subject, ev) + p.emit(ev) return } // Within the window: buffer the latest event and arm a trailing - // publish. If a timer is already armed, we just overwrite p.pending so - // the trailing publish carries the freshest data. + // emit. If a timer is already armed, we just overwrite p.pending so + // the trailing emit carries the freshest data. p.pending = &ev if p.timer == nil { - delay := p.interval - now.Sub(p.lastPublishedAt) + delay := p.interval - now.Sub(p.lastEmittedAt) p.timer = time.AfterFunc(delay, p.flushPending) } p.mu.Unlock() } -// flushPending is the trailing-edge publisher fired by the AfterFunc timer. -// It clears the pending slot under the lock, then publishes outside the -// lock so Publish never blocks an in-progress OnDownload call. -func (p *DebouncedInstallProgressPublisher) flushPending() { +// flushPending is the trailing-edge emitter fired by the AfterFunc timer. +// It clears the pending slot under the lock, then emits outside the lock so +// emit never blocks an in-progress OnDownload call. +func (p *DebouncedInstallProgressSink) flushPending() { p.mu.Lock() p.timer = nil pending := p.pending p.pending = nil if pending != nil { - p.lastPublishedAt = time.Now() + p.lastEmittedAt = time.Now() } p.mu.Unlock() if pending != nil { - _ = p.client.Publish(p.subject, *pending) + p.emit(*pending) } } -// Flush publishes any pending buffered event synchronously and stops the +// Flush emits any pending buffered event synchronously and stops the // pending timer. Safe to call multiple times. Callers MUST defer Flush -// after constructing the publisher so the terminal percentage reaches the +// after constructing the sink so the terminal percentage reaches the // master even on error returns. -func (p *DebouncedInstallProgressPublisher) Flush() { +func (p *DebouncedInstallProgressSink) Flush() { p.mu.Lock() if p.timer != nil { p.timer.Stop() @@ -115,6 +114,6 @@ func (p *DebouncedInstallProgressPublisher) Flush() { p.pending = nil p.mu.Unlock() if pending != nil { - _ = p.client.Publish(p.subject, *pending) + p.emit(*pending) } } diff --git a/core/services/nodes/install_progress_publisher_test.go b/core/services/nodes/install_progress_publisher_test.go index 04073cebe..03da1fbf4 100644 --- a/core/services/nodes/install_progress_publisher_test.go +++ b/core/services/nodes/install_progress_publisher_test.go @@ -1,48 +1,71 @@ package nodes import ( + "sync" "time" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" ) -var _ = Describe("DebouncedInstallProgressPublisher", func() { - It("publishes the first event immediately and debounces subsequent ones within the window", func() { - mc := newScriptedMessagingClient() - pub := NewDebouncedInstallProgressPublisher(mc, "n1", "op1", "vllm", 50*time.Millisecond) +// emitRecorder is the emit callback a carrier would hand the sink. The +// trailing debounce fires it from a timer goroutine, hence the lock. +type emitRecorder struct { + mu sync.Mutex + events []workerctl.BackendInstallProgressEvent +} + +func (r *emitRecorder) emit(ev workerctl.BackendInstallProgressEvent) { + r.mu.Lock() + defer r.mu.Unlock() + r.events = append(r.events, ev) +} + +func (r *emitRecorder) emitted() []workerctl.BackendInstallProgressEvent { + r.mu.Lock() + defer r.mu.Unlock() + return append([]workerctl.BackendInstallProgressEvent(nil), r.events...) +} + +var _ = Describe("DebouncedInstallProgressSink", func() { + It("emits the first event immediately and debounces subsequent ones within the window", func() { + rec := &emitRecorder{} + sink := NewDebouncedInstallProgressSink(rec.emit, "n1", "op1", "vllm", 50*time.Millisecond) // Three rapid-fire ticks within the debounce window. - pub.OnDownload("vllm.tar.zst", "100 MB", "1 GB", 10.0) - pub.OnDownload("vllm.tar.zst", "200 MB", "1 GB", 20.0) - pub.OnDownload("vllm.tar.zst", "300 MB", "1 GB", 30.0) - pub.Flush() + sink.OnDownload("vllm.tar.zst", "100 MB", "1 GB", 10.0) + sink.OnDownload("vllm.tar.zst", "200 MB", "1 GB", 20.0) + sink.OnDownload("vllm.tar.zst", "300 MB", "1 GB", 30.0) + sink.Flush() - // First event publishes immediately; the others coalesce; Flush guarantees a final. - // So we expect at least 2 publishes and at most 4 (lead + final + any window-bounded). - Eventually(func() int { - return len(mc.publishCalls(messaging.SubjectNodeBackendInstallProgress("n1", "op1"))) - }, "1s").Should(BeNumerically(">=", 2)) - calls := mc.publishCalls(messaging.SubjectNodeBackendInstallProgress("n1", "op1")) - Expect(len(calls)).To(BeNumerically("<=", 4), - "three ticks within the debounce window should produce at most ~4 publishes") + // First event emits immediately; the others coalesce; Flush guarantees a final. + // So we expect at least 2 emits and at most 4 (lead + final + any window-bounded). + Eventually(func() int { return len(rec.emitted()) }, "1s").Should(BeNumerically(">=", 2)) + Expect(len(rec.emitted())).To(BeNumerically("<=", 4), + "three ticks within the debounce window should produce at most ~4 emits") + for _, ev := range rec.emitted() { + Expect(ev.OpID).To(Equal("op1")) + Expect(ev.NodeID).To(Equal("n1")) + Expect(ev.Backend).To(Equal("vllm")) + Expect(ev.Phase).To(Equal(workerctl.PhaseDownloading)) + } }) - It("publishes the final event after Flush with the latest percentage", func() { - mc := newScriptedMessagingClient() - pub := NewDebouncedInstallProgressPublisher(mc, "n1", "op1", "vllm", 50*time.Millisecond) + It("emits the final event after Flush with the latest percentage", func() { + rec := &emitRecorder{} + sink := NewDebouncedInstallProgressSink(rec.emit, "n1", "op1", "vllm", 50*time.Millisecond) - pub.OnDownload("vllm.tar.zst", "1 GB", "1 GB", 100.0) - pub.Flush() + sink.OnDownload("vllm.tar.zst", "1 GB", "1 GB", 100.0) + sink.Flush() Eventually(func() float64 { - calls := mc.publishCalls(messaging.SubjectNodeBackendInstallProgress("n1", "op1")) - if len(calls) == 0 { + evs := rec.emitted() + if len(evs) == 0 { return -1 } - return calls[len(calls)-1].Percentage + return evs[len(evs)-1].Percentage }, "1s").Should(Equal(100.0)) }) }) diff --git a/core/services/nodes/interfaces.go b/core/services/nodes/interfaces.go index aafa0e47f..41e175e0b 100644 --- a/core/services/nodes/interfaces.go +++ b/core/services/nodes/interfaces.go @@ -2,14 +2,15 @@ package nodes import ( "context" + "net" "time" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" grpc "github.com/mudler/LocalAI/pkg/grpc" ) type ExactModelStopper interface { - StopModelReplica(ctx context.Context, nodeID string, replica NodeModel, force bool) (messaging.ModelStopReply, error) + StopModelReplica(ctx context.Context, nodeID string, replica NodeModel, force bool) (workerctl.ModelStopReply, error) } type ModelCleanupRegistry interface { @@ -146,9 +147,11 @@ type NodeManager interface { RemoveAllNodeModelReplicas(ctx context.Context, nodeID, modelName string) error } -// BackendClientFactory creates gRPC backend clients. +// BackendClientFactory creates gRPC backend clients. It takes the node id +// because a dialer that must know WHICH node it is reaching, as a tunnel does, +// cannot recover it from the address; a direct dialer ignores it. type BackendClientFactory interface { - NewClient(address string, parallel bool) grpc.Backend + NewClient(nodeID, address string, parallel bool) grpc.Backend } // tokenClientFactory is the default BackendClientFactory that creates gRPC @@ -157,9 +160,23 @@ type tokenClientFactory struct { token string } -func (f *tokenClientFactory) NewClient(address string, parallel bool) grpc.Backend { +func (f *tokenClientFactory) NewClient(_, address string, parallel bool) grpc.Backend { if f.token != "" { return grpc.NewClientWithToken(address, parallel, nil, false, f.token) } return grpc.NewClient(address, parallel, nil, false) } + +// WorkerNetDialerFor returns the dial function that reaches one worker's own +// HTTP server, in the shape http.Transport.DialContext and +// websocket.Dialer.NetDialContext take. It is keyed by node id, not address, +// because two workers can report the same HTTP address (NAT, loopback) and a +// tunnel must still reach the right one. +type WorkerNetDialerFor func(nodeID string) func(ctx context.Context, network, addr string) (net.Conn, error) + +// DirectWorkerNetDialer dials the address it is handed, whatever the node. +func DirectWorkerNetDialer() WorkerNetDialerFor { + // Aggressive keepalive suits the long LAN transfers the file stager makes. + dial := (&net.Dialer{Timeout: 30 * time.Second, KeepAlive: 15 * time.Second}).DialContext + return func(string) func(context.Context, string, string) (net.Conn, error) { return dial } +} diff --git a/core/services/nodes/managers_distributed.go b/core/services/nodes/managers_distributed.go index 4132eca79..0ea607cf0 100644 --- a/core/services/nodes/managers_distributed.go +++ b/core/services/nodes/managers_distributed.go @@ -10,11 +10,10 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/gallery" "github.com/mudler/LocalAI/core/services/galleryop" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/system" "github.com/mudler/xlog" - "github.com/nats-io/nats.go" ) // DistributedModelManager wraps a local ModelManager and adds NATS fan-out @@ -213,8 +212,8 @@ func (d *DistributedBackendManager) enqueueAndDrainBackendOp(ctx context.Context continue } - // Record failure for backoff. If it's an ErrNoResponders, the node's - // gone AWOL - mark unhealthy so the router stops picking it too. + // Record failure for backoff. If there is no route to the node, mark it + // unhealthy so the router stops picking it too. errMsg := applyErr.Error() // Worker-still-installing is a "soft" failure: the worker is most @@ -234,8 +233,8 @@ func (d *DistributedBackendManager) enqueueAndDrainBackendOp(ctx context.Context continue } - if errors.Is(applyErr, nats.ErrNoResponders) { - xlog.Warn("No NATS responders for node, marking unhealthy", "node", node.Name, "nodeID", node.ID) + if errors.Is(applyErr, ErrNoRoute) { + xlog.Warn("No route to node, marking unhealthy", "node", node.Name, "nodeID", node.ID) d.registry.MarkUnhealthy(ctx, node.ID) } if id, err := d.findPendingRow(ctx, node.ID, backend, op); err == nil { @@ -333,7 +332,7 @@ func (d *DistributedBackendManager) DeleteBackendDetailed(ctx context.Context, n // Pending/offline/draining nodes are skipped because they aren't expected to // answer NATS requests, and so are non-backend workers, which do not subscribe // to backend.list at all; unhealthy backend nodes are still queried — -// ErrNoResponders then marks them unhealthy and the loop continues. +// ErrNoRoute then marks them unhealthy and the loop continues. func (d *DistributedBackendManager) ListBackends() (gallery.SystemBackends, error) { result := make(gallery.SystemBackends) allNodes, err := d.registry.List(context.Background()) @@ -346,7 +345,7 @@ func (d *DistributedBackendManager) ListBackends() (gallery.SystemBackends, erro continue } // Only backend workers subscribe to backend.list. Asking an agent - // worker can only answer "no responders", which the error handling + // worker can only answer "no route", which the error handling // below reads as a node that has gone away, so every poll of this view // marked every agent node unhealthy and its next heartbeat marked it // healthy again. The backend-op fan-out skips them for the same reason. @@ -355,8 +354,8 @@ func (d *DistributedBackendManager) ListBackends() (gallery.SystemBackends, erro } reply, err := d.adapter.ListBackends(node.ID) if err != nil { - if errors.Is(err, nats.ErrNoResponders) { - xlog.Warn("No NATS responders for node, marking unhealthy", "node", node.Name, "nodeID", node.ID) + if errors.Is(err, ErrNoRoute) { + xlog.Warn("No route to node, marking unhealthy", "node", node.Name, "nodeID", node.ID) d.registry.MarkUnhealthy(context.Background(), node.ID) continue } @@ -476,7 +475,7 @@ func (d *DistributedBackendManager) InstallBackend(ctx context.Context, op *gall // per-node sink (so OpStatus.Nodes gets a "downloading" tick // per file/percentage with node attribution). Defined inside the // loop so each node captures its own node.Name into the closure. - onProgress := func(ev messaging.BackendInstallProgressEvent) { + onProgress := func(ev workerctl.BackendInstallProgressEvent) { if progressCb != nil { progressCb(ev.FileName, ev.Current, ev.Total, ev.Percentage) } @@ -496,7 +495,7 @@ func (d *DistributedBackendManager) InstallBackend(ctx context.Context, op *gall // nil-callback shortcut: when there is nothing to deliver to, // hand the adapter a nil onProgress so it skips the per-op NATS // subscription. Matches the pre-Phase-4 bridgeProgressCb semantics. - var onProgressArg func(messaging.BackendInstallProgressEvent) + var onProgressArg func(workerctl.BackendInstallProgressEvent) if progressCb != nil || d.progressSink != nil { onProgressArg = onProgress } @@ -538,7 +537,7 @@ func (d *DistributedBackendManager) InstallBackend(ctx context.Context, op *gall // worker has no platform variant for a linux-only backend) and leaves a // forever-retrying pending_backend_ops row. // -// Rolling-update fallback: when a worker returns nats.ErrNoResponders on +// Rolling-update fallback: when a worker returns ErrNoRoute on // backend.upgrade, we try the legacy backend.install Force=true path so a // new master + old worker still converges. Drop the fallback once every // worker in the fleet is on 2026-05-08 or newer. @@ -576,7 +575,7 @@ func (d *DistributedBackendManager) UpgradeBackend(ctx context.Context, op *gall // InstallBackend does. Defined per-node so each closure captures its own // node.Name. Without this an upgrade blocks opaque at progress 0 for the // whole 15m round-trip (the original "reinstalling but nothing happens"). - onProgress := func(ev messaging.BackendInstallProgressEvent) { + onProgress := func(ev workerctl.BackendInstallProgressEvent) { if progressCb != nil { progressCb(ev.FileName, ev.Current, ev.Total, ev.Percentage) } @@ -593,7 +592,7 @@ func (d *DistributedBackendManager) UpgradeBackend(ctx context.Context, op *gall }) } } - var onProgressArg func(messaging.BackendInstallProgressEvent) + var onProgressArg func(workerctl.BackendInstallProgressEvent) if progressCb != nil || d.progressSink != nil { onProgressArg = onProgress } @@ -601,7 +600,7 @@ func (d *DistributedBackendManager) UpgradeBackend(ctx context.Context, op *gall if err != nil { // Rolling-update fallback: an older worker doesn't know // backend.upgrade. Try the legacy install-with-force path. - if errors.Is(err, nats.ErrNoResponders) { + if errors.Is(err, ErrNoRoute) { instReply, instErr := d.adapter.installWithForceFallback(node.ID, name, string(galleriesJSON), "", "", "", 0, opID, onProgressArg) if instErr != nil { return instErr diff --git a/core/services/nodes/managers_distributed_test.go b/core/services/nodes/managers_distributed_test.go index b83200eeb..490f3e098 100644 --- a/core/services/nodes/managers_distributed_test.go +++ b/core/services/nodes/managers_distributed_test.go @@ -18,6 +18,7 @@ import ( "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" ) // scriptedMessagingClient maps a NATS subject to a canned reply payload @@ -29,22 +30,10 @@ type scriptedMessagingClient struct { errs map[string]error calls []requestCall matchedReplies map[string][]matchedReply - publishes []progressPublishCall scheduledProgressPublishes []scheduledProgressPublish subscribes []string } -// progressPublishCall records a single Publish invocation. The progress -// publisher tests assert on the sequence of BackendInstallProgressEvent -// values written to a per-op subject, so we capture both subject and the -// decoded event. Named to avoid clashing with the simpler `publishCall` -// already defined in unloader_test.go (which stores raw JSON bytes for -// non-progress assertions). -type progressPublishCall struct { - Subject string - Event messaging.BackendInstallProgressEvent -} - // scheduledProgressPublish queues a batch of BackendInstallProgressEvent // values to be delivered the next time Subscribe is called with the matching // subject. This lets master-side tests assert that the adapter installs its @@ -52,7 +41,7 @@ type progressPublishCall struct { // delivered as soon as the subscription appears. type scheduledProgressPublish struct { subject string - events []messaging.BackendInstallProgressEvent + events []workerctl.BackendInstallProgressEvent } // matchedReply lets a test script a canned reply that only fires when the @@ -60,7 +49,7 @@ type scheduledProgressPublish struct { // distinguish "install Force=true" (the fallback) from "install Force=false" // on the same subject. type matchedReply struct { - pred func(messaging.BackendInstallRequest) bool + pred func(workerctl.BackendInstallRequest) bool reply []byte fallback []byte fallbackErr error @@ -106,7 +95,7 @@ func (s *scriptedMessagingClient) scriptNoResponders(subject string) { // If `pred` returns false (or the unmarshal of the payload into the // predicate's expected type fails), the subject falls through to whatever // was scripted before (or to the unscripted default ErrNoResponders). -func (s *scriptedMessagingClient) scriptReplyMatching(subject string, pred func(messaging.BackendInstallRequest) bool, reply messaging.BackendInstallReply) { +func (s *scriptedMessagingClient) scriptReplyMatching(subject string, pred func(workerctl.BackendInstallRequest) bool, reply workerctl.BackendInstallReply) { raw, err := json.Marshal(reply) Expect(err).ToNot(HaveOccurred()) s.mu.Lock() @@ -131,7 +120,7 @@ func (s *scriptedMessagingClient) Request(subject string, data []byte, timeout t // Predicate-matched replies take precedence over flat scriptReply. if matchers, ok := s.matchedReplies[subject]; ok { - var req messaging.BackendInstallRequest + var req workerctl.BackendInstallRequest _ = json.Unmarshal(data, &req) for _, m := range matchers { if m.pred(req) { @@ -161,45 +150,17 @@ func (s *scriptedMessagingClient) Request(subject string, data []byte, timeout t return nil, &fakeNoRespondersErr{} } -// Publish records each call so progress-publisher tests can assert on the -// stream of events written to a subject. The real messaging.Client JSON -// encodes the payload before sending, but our publisher hands a typed -// struct directly, so we handle both shapes. -func (s *scriptedMessagingClient) Publish(subject string, data any) error { - s.mu.Lock() - defer s.mu.Unlock() - switch ev := data.(type) { - case messaging.BackendInstallProgressEvent: - s.publishes = append(s.publishes, progressPublishCall{Subject: subject, Event: ev}) - case []byte: - var e messaging.BackendInstallProgressEvent - _ = json.Unmarshal(ev, &e) - s.publishes = append(s.publishes, progressPublishCall{Subject: subject, Event: e}) - } +// Publish drops every event: no spec reads what the frontend publishes +// through this fake. +func (s *scriptedMessagingClient) Publish(string, any) error { return nil } -// publishCalls returns every BackendInstallProgressEvent that was published -// to `subject`, in order. Lets tests assert on debounce behavior without -// depending on internal Publish timing. -func (s *scriptedMessagingClient) publishCalls(subject string) []messaging.BackendInstallProgressEvent { - s.mu.Lock() - defer s.mu.Unlock() - out := make([]messaging.BackendInstallProgressEvent, 0) - for _, c := range s.publishes { - if c.Subject != subject { - continue - } - out = append(out, c.Event) - } - return out -} - // scheduleProgressPublish queues a set of BackendInstallProgressEvent values // to be delivered on the next Subscribe call matching the per-op progress // subject. A short delay before delivery gives the subscriber time to install // its message handler before the events arrive. -func (s *scriptedMessagingClient) scheduleProgressPublish(nodeID, opID string, events []messaging.BackendInstallProgressEvent) { +func (s *scriptedMessagingClient) scheduleProgressPublish(nodeID, opID string, events []workerctl.BackendInstallProgressEvent) { s.mu.Lock() defer s.mu.Unlock() s.scheduledProgressPublishes = append(s.scheduledProgressPublishes, scheduledProgressPublish{ @@ -386,9 +347,9 @@ var _ = Describe("DistributedBackendManager", func() { n2 := registerHealthyBackend("worker-b", "10.0.0.2:50051") mc.scriptReply(messaging.SubjectNodeBackendInstall(n1.ID), - messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) + workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) mc.scriptReply(messaging.SubjectNodeBackendInstall(n2.ID), - messaging.BackendInstallReply{Success: true, Address: "10.0.0.2:50100"}) + workerctl.BackendInstallReply{Success: true, Address: "10.0.0.2:50100"}) Expect(mgr.InstallBackend(ctx, op("vllm-development"), nil)).To(Succeed()) }) @@ -400,9 +361,9 @@ var _ = Describe("DistributedBackendManager", func() { n2 := registerHealthyBackend("nvidia-thor", "10.0.0.2:50051") mc.scriptReply(messaging.SubjectNodeBackendInstall(n1.ID), - messaging.BackendInstallReply{Success: false, Error: "no child with platform linux/arm64 in index quay.io/...master-cpu-vllm"}) + workerctl.BackendInstallReply{Success: false, Error: "no child with platform linux/arm64 in index quay.io/...master-cpu-vllm"}) mc.scriptReply(messaging.SubjectNodeBackendInstall(n2.ID), - messaging.BackendInstallReply{Success: false, Error: "disk full"}) + workerctl.BackendInstallReply{Success: false, Error: "disk full"}) err := mgr.InstallBackend(ctx, op("vllm-development"), nil) Expect(err).To(HaveOccurred()) @@ -420,9 +381,9 @@ var _ = Describe("DistributedBackendManager", func() { bad := registerHealthyBackend("worker-bad", "10.0.0.2:50051") mc.scriptReply(messaging.SubjectNodeBackendInstall(ok.ID), - messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) + workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) mc.scriptReply(messaging.SubjectNodeBackendInstall(bad.ID), - messaging.BackendInstallReply{Success: false, Error: "out of memory"}) + workerctl.BackendInstallReply{Success: false, Error: "out of memory"}) err := mgr.InstallBackend(ctx, op("vllm-development"), nil) Expect(err).To(HaveOccurred()) @@ -459,7 +420,7 @@ var _ = Describe("DistributedBackendManager", func() { other := registerHealthyBackend("worker-other", "10.0.0.2:50051") mc.scriptReply(messaging.SubjectNodeBackendInstall(target.ID), - messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) + workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) // No reply scripted for `other`: if InstallBackend fans out // to it, the fakeNoRespondersErr default would surface and // the test would fail. @@ -545,8 +506,8 @@ var _ = Describe("DistributedBackendManager", func() { // The worker finished installing in the background. Script // backend.list on the same scriptedMessagingClient so the // manager's ListBackends fan-out reports the backend. - mc.scriptReply(messaging.SubjectNodeBackendList(node.ID), messaging.BackendListReply{ - Backends: []messaging.NodeBackendInfo{{Name: "vllm"}}, + mc.scriptReply(messaging.SubjectNodeBackendList(node.ID), workerctl.BackendListReply{ + Backends: []workerctl.NodeBackendInfo{{Name: "vllm"}}, }) backends, listErr := mgr.ListBackends() @@ -581,8 +542,8 @@ var _ = Describe("DistributedBackendManager", func() { // Worker finishes installing in the background. backend.list now // confirms presence; ListBackends should proactively clear the row. - mc.scriptReply(messaging.SubjectNodeBackendList(node.ID), messaging.BackendListReply{ - Backends: []messaging.NodeBackendInfo{{Name: "vllm"}}, + mc.scriptReply(messaging.SubjectNodeBackendList(node.ID), workerctl.BackendListReply{ + Backends: []workerctl.NodeBackendInfo{{Name: "vllm"}}, }) backends, listErr := mgr.ListBackends() @@ -599,8 +560,8 @@ var _ = Describe("DistributedBackendManager", func() { Expect(registry.UpsertPendingBackendOp(ctx, node.ID, "vllm", OpBackendUpgrade, []byte("[]"))).To(Succeed()) - mc.scriptReply(messaging.SubjectNodeBackendList(node.ID), messaging.BackendListReply{ - Backends: []messaging.NodeBackendInfo{{Name: "vllm"}}, + mc.scriptReply(messaging.SubjectNodeBackendList(node.ID), workerctl.BackendListReply{ + Backends: []workerctl.NodeBackendInfo{{Name: "vllm"}}, }) _, listErr := mgr.ListBackends() @@ -615,8 +576,8 @@ var _ = Describe("DistributedBackendManager", func() { It("invokes progressCb once per worker-published progress event", func() { node := registerHealthyBackend("worker-prog", "10.0.0.7:50051") - mc.scriptReply(messaging.SubjectNodeBackendInstall(node.ID), messaging.BackendInstallReply{Success: true, Address: "10.0.0.7:50051"}) - mc.scheduleProgressPublish(node.ID, "op-prog-1", []messaging.BackendInstallProgressEvent{ + mc.scriptReply(messaging.SubjectNodeBackendInstall(node.ID), workerctl.BackendInstallReply{Success: true, Address: "10.0.0.7:50051"}) + mc.scheduleProgressPublish(node.ID, "op-prog-1", []workerctl.BackendInstallProgressEvent{ {OpID: "op-prog-1", NodeID: node.ID, Backend: "vllm", FileName: "vllm.tar", Current: "100 MB", Total: "1 GB", Percentage: 10}, {OpID: "op-prog-1", NodeID: node.ID, Backend: "vllm", FileName: "vllm.tar", Current: "1 GB", Total: "1 GB", Percentage: 100}, }) @@ -659,7 +620,7 @@ var _ = Describe("DistributedBackendManager", func() { Context("InstallBackend tolerates silent (pre-Phase-2) workers", func() { It("completes successfully even when no progress events are ever published", func() { node := registerHealthyBackend("worker-silent", "10.0.0.8:50051") - mc.scriptReply(messaging.SubjectNodeBackendInstall(node.ID), messaging.BackendInstallReply{Success: true, Address: "10.0.0.8:50051"}) + mc.scriptReply(messaging.SubjectNodeBackendInstall(node.ID), workerctl.BackendInstallReply{Success: true, Address: "10.0.0.8:50051"}) // NO scheduleProgressPublish call - silent worker. var ticks int @@ -702,7 +663,7 @@ var _ = Describe("DistributedBackendManager", func() { It("emits a success entry for each healthy node visited", func() { node := registerHealthyBackend("worker-ok", "10.0.0.9:50051") mc.scriptReply(messaging.SubjectNodeBackendInstall(node.ID), - messaging.BackendInstallReply{Success: true, Address: "10.0.0.9:50051"}) + workerctl.BackendInstallReply{Success: true, Address: "10.0.0.9:50051"}) opVal := op("vllm") opVal.ID = "op-node-success" @@ -731,9 +692,9 @@ var _ = Describe("DistributedBackendManager", func() { It("emits downloading entries from progress events", func() { node := registerHealthyBackend("worker-dl", "10.0.0.11:50051") mc.scriptReply(messaging.SubjectNodeBackendInstall(node.ID), - messaging.BackendInstallReply{Success: true}) - mc.scheduleProgressPublish(node.ID, "op-node-dl", []messaging.BackendInstallProgressEvent{ - {OpID: "op-node-dl", NodeID: node.ID, Backend: "vllm", FileName: "vllm.tar", Current: "1 GB", Total: "1 GB", Percentage: 100, Phase: messaging.PhaseDownloading}, + workerctl.BackendInstallReply{Success: true}) + mc.scheduleProgressPublish(node.ID, "op-node-dl", []workerctl.BackendInstallProgressEvent{ + {OpID: "op-node-dl", NodeID: node.ID, Backend: "vllm", FileName: "vllm.tar", Current: "1 GB", Total: "1 GB", Percentage: 100, Phase: workerctl.PhaseDownloading}, }) opVal := op("vllm") @@ -766,13 +727,13 @@ var _ = Describe("DistributedBackendManager", func() { scriptInstalled := func(backend string, nodeIDs ...string) { for _, id := range nodeIDs { mc.scriptReply(messaging.SubjectNodeBackendList(id), - messaging.BackendListReply{Backends: []messaging.NodeBackendInfo{{Name: backend}}}) + workerctl.BackendListReply{Backends: []workerctl.NodeBackendInfo{{Name: backend}}}) } } scriptNoBackends := func(nodeIDs ...string) { for _, id := range nodeIDs { mc.scriptReply(messaging.SubjectNodeBackendList(id), - messaging.BackendListReply{Backends: nil}) + workerctl.BackendListReply{Backends: nil}) } } @@ -783,9 +744,9 @@ var _ = Describe("DistributedBackendManager", func() { scriptInstalled("vllm-development", n1.ID, n2.ID) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(n1.ID), - messaging.BackendUpgradeReply{Success: false, Error: "image manifest not found"}) + workerctl.BackendUpgradeReply{Success: false, Error: "image manifest not found"}) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(n2.ID), - messaging.BackendUpgradeReply{Success: false, Error: "registry unauthorized"}) + workerctl.BackendUpgradeReply{Success: false, Error: "registry unauthorized"}) err := mgr.UpgradeBackend(ctx, upgradeOp("vllm-development"), nil) Expect(err).To(HaveOccurred()) @@ -801,7 +762,7 @@ var _ = Describe("DistributedBackendManager", func() { n1 := registerHealthyBackend("worker-a", "10.0.0.1:50051") scriptInstalled("vllm-development", n1.ID) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(n1.ID), - messaging.BackendUpgradeReply{Success: true}) + workerctl.BackendUpgradeReply{Success: true}) Expect(mgr.UpgradeBackend(ctx, upgradeOp("vllm-development"), nil)).To(Succeed()) }) }) @@ -819,7 +780,7 @@ var _ = Describe("DistributedBackendManager", func() { scriptInstalled("cpu-insightface-development", has.ID) scriptNoBackends(lacks.ID) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(has.ID), - messaging.BackendUpgradeReply{Success: true}) + workerctl.BackendUpgradeReply{Success: true}) // Deliberately don't script SubjectNodeBackendUpgrade for `lacks`: // if the manager attempts it, the scripted-client default returns // fakeNoRespondersErr and the assertion below fails loudly. @@ -847,9 +808,9 @@ var _ = Describe("DistributedBackendManager", func() { scriptInstalled("vllm-development", n1.ID, n2.ID) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(n1.ID), - messaging.BackendUpgradeReply{Success: true}) + workerctl.BackendUpgradeReply{Success: true}) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(n2.ID), - messaging.BackendUpgradeReply{Success: true}) + workerctl.BackendUpgradeReply{Success: true}) op := upgradeOp("vllm-development") op.TargetNodeID = n2.ID @@ -877,7 +838,7 @@ var _ = Describe("DistributedBackendManager", func() { scriptInstalled("vllm-development", has.ID) scriptNoBackends(lacks.ID) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(has.ID), - messaging.BackendUpgradeReply{Success: true}) + workerctl.BackendUpgradeReply{Success: true}) op := upgradeOp("vllm-development") op.TargetNodeID = lacks.ID @@ -915,12 +876,12 @@ var _ = Describe("DistributedBackendManager", func() { }) // Rolling-update fallback: pre-2026-05-08 workers don't subscribe to - // backend.upgrade, so the manager catches nats.ErrNoResponders and + // backend.upgrade, so the adapter reports ErrNoRoute and the manager // re-fires the legacy backend.install Force=true on the same node. // Drop these specs once the fallback path itself is removed (see // managers_distributed.go UpgradeBackend godoc for the deprecation). Context("rolling-update fallback", func() { - It("falls back to backend.install Force=true when upgrade returns ErrNoResponders", func() { + It("falls back to backend.install Force=true when upgrade returns ErrNoRoute", func() { n := registerHealthyBackend("worker-old", "10.0.0.1:50051") scriptInstalled("vllm-development", n.ID) @@ -928,18 +889,18 @@ var _ = Describe("DistributedBackendManager", func() { mc.scriptNoResponders(messaging.SubjectNodeBackendUpgrade(n.ID)) // Fallback re-fires legacy backend.install with Force=true. mc.scriptReplyMatching(messaging.SubjectNodeBackendInstall(n.ID), - func(req messaging.BackendInstallRequest) bool { return req.Force }, - messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) + func(req workerctl.BackendInstallRequest) bool { return req.Force }, + workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) Expect(mgr.UpgradeBackend(ctx, upgradeOp("vllm-development"), nil)).To(Succeed()) }) - It("returns the upgrade error when it is not ErrNoResponders", func() { + It("returns the upgrade error when it is not ErrNoRoute", func() { n := registerHealthyBackend("worker-bad", "10.0.0.1:50051") scriptInstalled("vllm-development", n.ID) mc.scriptReply(messaging.SubjectNodeBackendUpgrade(n.ID), - messaging.BackendUpgradeReply{Success: false, Error: "disk full"}) + workerctl.BackendUpgradeReply{Success: false, Error: "disk full"}) err := mgr.UpgradeBackend(ctx, upgradeOp("vllm-development"), nil) Expect(err).To(HaveOccurred()) @@ -955,9 +916,9 @@ var _ = Describe("DistributedBackendManager", func() { n2 := registerHealthyBackend("worker-b", "10.0.0.2:50051") mc.scriptReply(messaging.SubjectNodeBackendDelete(n1.ID), - messaging.BackendDeleteReply{Success: false, Error: "backend not installed"}) + workerctl.BackendDeleteReply{Success: false, Error: "backend not installed"}) mc.scriptReply(messaging.SubjectNodeBackendDelete(n2.ID), - messaging.BackendDeleteReply{Success: false, Error: "permission denied"}) + workerctl.BackendDeleteReply{Success: false, Error: "permission denied"}) err := mgr.DeleteBackend("vllm-development") Expect(err).To(HaveOccurred()) @@ -972,7 +933,7 @@ var _ = Describe("DistributedBackendManager", func() { It("returns nil", func() { n1 := registerHealthyBackend("worker-a", "10.0.0.1:50051") mc.scriptReply(messaging.SubjectNodeBackendDelete(n1.ID), - messaging.BackendDeleteReply{Success: true}) + workerctl.BackendDeleteReply{Success: true}) Expect(mgr.DeleteBackend("vllm-development")).To(Succeed()) }) }) diff --git a/core/services/nodes/model_cleanup_test.go b/core/services/nodes/model_cleanup_test.go index c09bbef4d..0d28cc535 100644 --- a/core/services/nodes/model_cleanup_test.go +++ b/core/services/nodes/model_cleanup_test.go @@ -6,7 +6,7 @@ import ( "sync" "time" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -47,7 +47,7 @@ func (f *fakeCleanupRegistry) RecordModelCleanupFailure(_ context.Context, _, _ type fakeExactStopper struct { mu sync.Mutex - replies []messaging.ModelStopReply + replies []workerctl.ModelStopReply errs []error calls []NodeModel block chan struct{} @@ -84,7 +84,7 @@ type blockingExactStopper struct { calls int } -func (f *blockingExactStopper) StopModelReplica(_ context.Context, _ string, _ NodeModel, _ bool) (messaging.ModelStopReply, error) { +func (f *blockingExactStopper) StopModelReplica(_ context.Context, _ string, _ NodeModel, _ bool) (workerctl.ModelStopReply, error) { f.mu.Lock() f.calls++ if f.calls == 1 { @@ -92,10 +92,10 @@ func (f *blockingExactStopper) StopModelReplica(_ context.Context, _ string, _ N } f.mu.Unlock() <-f.release - return messaging.ModelStopReply{Matched: true, Terminated: true}, nil + return workerctl.ModelStopReply{Matched: true, Terminated: true}, nil } -func (f *fakeExactStopper) StopModelReplica(_ context.Context, _ string, replica NodeModel, _ bool) (messaging.ModelStopReply, error) { +func (f *fakeExactStopper) StopModelReplica(_ context.Context, _ string, replica NodeModel, _ bool) (workerctl.ModelStopReply, error) { if f.block != nil { <-f.block } @@ -103,7 +103,7 @@ func (f *fakeExactStopper) StopModelReplica(_ context.Context, _ string, replica defer f.mu.Unlock() i := len(f.calls) f.calls = append(f.calls, replica) - var reply messaging.ModelStopReply + var reply workerctl.ModelStopReply var err error if i < len(f.replies) { reply = f.replies[i] @@ -120,7 +120,7 @@ var _ = Describe("ModelCleanupService", func() { It("deletes only replicas whose termination is confirmed", func() { registry := &fakeCleanupRegistry{} - stopper := &fakeExactStopper{replies: []messaging.ModelStopReply{{Matched: true, Terminated: true}}} + stopper := &fakeExactStopper{replies: []workerctl.ModelStopReply{{Matched: true, Terminated: true}}} service := NewModelCleanupService(registry, stopper) service.now = func() time.Time { return now } service.Cleanup(context.Background(), []NodeModel{{NodeID: "n1", ModelName: "m", ReplicaIndex: 3}}, false) @@ -130,7 +130,7 @@ var _ = Describe("ModelCleanupService", func() { It("treats exact process absence as idempotent success", func() { registry := &fakeCleanupRegistry{} - stopper := &fakeExactStopper{replies: []messaging.ModelStopReply{{Matched: false, Terminated: true}}} + stopper := &fakeExactStopper{replies: []workerctl.ModelStopReply{{Matched: false, Terminated: true}}} service := NewModelCleanupService(registry, stopper) service.Cleanup(context.Background(), []NodeModel{{NodeID: "n1", ModelName: "m"}}, false) Expect(registry.removed).To(HaveLen(1)) @@ -149,7 +149,7 @@ var _ = Describe("ModelCleanupService", func() { It("retries transient failures and later removes the row", func() { registry := &fakeCleanupRegistry{} - stopper := &fakeExactStopper{errs: []error{errors.New("timeout"), nil}, replies: []messaging.ModelStopReply{{}, {Matched: true, Terminated: true}}} + stopper := &fakeExactStopper{errs: []error{errors.New("timeout"), nil}, replies: []workerctl.ModelStopReply{{}, {Matched: true, Terminated: true}}} service := NewModelCleanupService(registry, stopper) r := NodeModel{NodeID: "n1", ModelName: "m"} service.Cleanup(context.Background(), []NodeModel{r}, false) @@ -160,7 +160,7 @@ var _ = Describe("ModelCleanupService", func() { It("records a negative reply and tolerates a concurrent row deletion", func() { registry := &fakeCleanupRegistry{} - stopper := &fakeExactStopper{replies: []messaging.ModelStopReply{{Matched: true, Terminated: false, Error: "address mismatch"}}} + stopper := &fakeExactStopper{replies: []workerctl.ModelStopReply{{Matched: true, Terminated: false, Error: "address mismatch"}}} service := NewModelCleanupService(registry, stopper) service.Cleanup(context.Background(), []NodeModel{{NodeID: "n1", ModelName: "m"}}, false) Expect(registry.failures).To(Equal([]string{"address mismatch"})) @@ -168,7 +168,7 @@ var _ = Describe("ModelCleanupService", func() { It("leases due work so two runners do not own the same replica", func() { registry := &fakeCleanupRegistry{due: []NodeModel{{NodeID: "n1", ModelName: "m"}}} - stopper := &fakeExactStopper{replies: []messaging.ModelStopReply{{Matched: true, Terminated: true}}} + stopper := &fakeExactStopper{replies: []workerctl.ModelStopReply{{Matched: true, Terminated: true}}} a := NewModelCleanupService(registry, stopper) b := NewModelCleanupService(registry, stopper) a.runOnce(context.Background()) diff --git a/core/services/nodes/noroute_reactions_test.go b/core/services/nodes/noroute_reactions_test.go new file mode 100644 index 000000000..3d35b725f --- /dev/null +++ b/core/services/nodes/noroute_reactions_test.go @@ -0,0 +1,161 @@ +package nodes + +import ( + "context" + "encoding/json" + "runtime" + "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/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// These specs pin the reactions the seams contract allows on ErrNoRoute: the +// status-only MarkUnhealthy, and the legacy install fallback for an upgrade. +// Each one drives the real caller with a scripted no-responders reply and reads +// the outcome back from the registry or the recorded requests, so a change to +// the reaction fails here instead of silently widening or dropping it. +var _ = Describe("ErrNoRoute reactions", func() { + var ( + registry *NodeRegistry + mc *scriptedMessagingClient + adapter *RemoteUnloaderAdapter + 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()) + mc = newScriptedMessagingClient() + adapter = NewRemoteUnloaderAdapter(nil, mc, 3*time.Minute, 15*time.Minute) + ctx = context.Background() + }) + + registerHealthy := func(name string) *BackendNode { + node := &BackendNode{Name: name, NodeType: NodeTypeBackend, Address: name + ":50051"} + Expect(registry.Register(ctx, node, true)).To(Succeed()) + fetched, err := registry.Get(ctx, node.ID) + Expect(err).ToNot(HaveOccurred()) + Expect(fetched.Status).To(Equal(StatusHealthy)) + return fetched + } + + statusOf := func(nodeID string) string { + n, err := registry.Get(ctx, nodeID) + Expect(err).ToNot(HaveOccurred()) + return n.Status + } + + pendingRow := func(nodeID, op string) (PendingBackendOp, bool) { + var rows []PendingBackendOp + Expect(registry.db.WithContext(ctx).Where("node_id = ? AND op = ?", nodeID, op).Find(&rows).Error).To(Succeed()) + if len(rows) == 0 { + return PendingBackendOp{}, false + } + return rows[0], true + } + + Describe("reconciler pending-op drain", func() { + var rc *ReplicaReconciler + + BeforeEach(func() { + rc = NewReplicaReconciler(ReplicaReconcilerOptions{ + Registry: registry, + Adapter: adapter, + DB: registry.db, + }) + }) + + It("falls back to the legacy forced install when the upgrade has no route", func() { + n := registerHealthy("worker-old") + Expect(registry.UpsertPendingBackendOp(ctx, n.ID, "vllm", OpBackendUpgrade, []byte("[]"))).To(Succeed()) + + mc.scriptNoResponders(messaging.SubjectNodeBackendUpgrade(n.ID)) + mc.scriptReplyMatching(messaging.SubjectNodeBackendInstall(n.ID), + func(req workerctl.BackendInstallRequest) bool { return req.Force }, + workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}) + + rc.drainPendingBackendOps(ctx) + + var forcedInstalls int + mc.mu.Lock() + for _, call := range mc.calls { + if call.Subject != messaging.SubjectNodeBackendInstall(n.ID) { + continue + } + var req workerctl.BackendInstallRequest + Expect(json.Unmarshal(call.Data, &req)).To(Succeed()) + if req.Force && req.Backend == "vllm" { + forcedInstalls++ + } + } + mc.mu.Unlock() + Expect(forcedInstalls).To(Equal(1)) + + _, stillQueued := pendingRow(n.ID, OpBackendUpgrade) + Expect(stillQueued).To(BeFalse(), "a successful fallback drains the row") + Expect(statusOf(n.ID)).To(Equal(StatusHealthy), "an old worker that answered the fallback is not unhealthy") + }) + + It("marks the node unhealthy when an op has no route, and still counts the attempt", func() { + n := registerHealthy("worker-gone") + Expect(registry.UpsertPendingBackendOp(ctx, n.ID, "vllm", OpBackendDelete, nil)).To(Succeed()) + mc.scriptNoResponders(messaging.SubjectNodeBackendDelete(n.ID)) + + rc.drainPendingBackendOps(ctx) + + Expect(statusOf(n.ID)).To(Equal(StatusUnhealthy)) + row, queued := pendingRow(n.ID, OpBackendDelete) + Expect(queued).To(BeTrue()) + Expect(row.Attempts).To(Equal(1)) + }) + + It("leaves the node healthy when the op times out", func() { + n := registerHealthy("worker-slow") + Expect(registry.UpsertPendingBackendOp(ctx, n.ID, "vllm", OpBackendDelete, nil)).To(Succeed()) + mc.scriptErr(messaging.SubjectNodeBackendDelete(n.ID), nats.ErrTimeout) + + rc.drainPendingBackendOps(ctx) + + Expect(statusOf(n.ID)).To(Equal(StatusHealthy)) + row, queued := pendingRow(n.ID, OpBackendDelete) + Expect(queued).To(BeTrue()) + Expect(row.Attempts).To(Equal(1)) + }) + }) + + Describe("DistributedBackendManager fan-out", func() { + var mgr *DistributedBackendManager + + BeforeEach(func() { + mgr = &DistributedBackendManager{ + local: stubLocalBackendManager{}, + adapter: adapter, + registry: registry, + } + }) + + It("marks a node with no route unhealthy and leaves an answering node healthy", func() { + gone := registerHealthy("worker-gone") + answering := registerHealthy("worker-answering") + mc.scriptNoResponders(messaging.SubjectNodeBackendDelete(gone.ID)) + mc.scriptReply(messaging.SubjectNodeBackendDelete(answering.ID), + workerctl.BackendDeleteReply{Success: false, Error: "backend not installed"}) + + Expect(mgr.DeleteBackend("vllm")).ToNot(Succeed()) + + Expect(statusOf(gone.ID)).To(Equal(StatusUnhealthy)) + Expect(statusOf(answering.ID)).To(Equal(StatusHealthy)) + }) + }) +}) diff --git a/core/services/nodes/pending_op_cleanup_test.go b/core/services/nodes/pending_op_cleanup_test.go index ad8610cc4..26c5443db 100644 --- a/core/services/nodes/pending_op_cleanup_test.go +++ b/core/services/nodes/pending_op_cleanup_test.go @@ -86,7 +86,7 @@ var _ = Describe("DeleteStalePendingBackendOps", func() { }) It("clears ops behind an unhealthy node with a stale heartbeat (never ages to offline)", func() { - // A node marked unhealthy on a NATS ErrNoResponders never transitions to + // A node marked unhealthy on an ErrNoRoute never transitions to // offline, so its ops must be reaped via the same stale-heartbeat path. sick := registerBackend("agx-orin-sick", "10.0.0.7:50051") Expect(registry.UpsertPendingBackendOp(ctx, sick, "llama-cpp-development", OpBackendUpgrade, nil)).To(Succeed()) diff --git a/core/services/nodes/reconciler.go b/core/services/nodes/reconciler.go index a14a3fa70..e1d7a09d3 100644 --- a/core/services/nodes/reconciler.go +++ b/core/services/nodes/reconciler.go @@ -10,11 +10,9 @@ import ( "time" "github.com/mudler/LocalAI/core/services/advisorylock" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes/prefixcache" - grpcclient "github.com/mudler/LocalAI/pkg/grpc" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/xlog" - "github.com/nats-io/nats.go" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" "gorm.io/gorm" @@ -41,7 +39,7 @@ const ( // Defaulted to a gRPC health probe but overridable for tests so we don't // need to stand up a real server. type ModelProber interface { - Probe(ctx context.Context, address string) ProbeOutcome + Probe(ctx context.Context, nodeID, address string) ProbeOutcome } // NodeProcessLister asks a worker which model backend processes it currently @@ -52,7 +50,7 @@ type ModelProber interface { // against the backend's own serving port cannot make that distinction, which // is why it is only the fallback for workers that do not answer. type NodeProcessLister interface { - ListRunningModels(nodeID string) (*messaging.ModelsRunningReply, error) + ListRunningModels(nodeID string) (*workerctl.ModelsRunningReply, error) } // probeTimeout bounds a single liveness probe. Kept short because a healthy @@ -62,10 +60,10 @@ type NodeProcessLister interface { const probeTimeout = 1 * time.Second // grpcModelProber does a short HealthCheck on the model's stored gRPC address. -type grpcModelProber struct{ token string } +type grpcModelProber struct{ clients BackendClientFactory } -func (g grpcModelProber) Probe(ctx context.Context, address string) ProbeOutcome { - client := grpcclient.NewClientWithToken(address, false, nil, false, g.token) +func (g grpcModelProber) Probe(ctx context.Context, nodeID, address string) ProbeOutcome { + client := g.clients.NewClient(nodeID, address, false) probeCtx, cancel := context.WithTimeout(ctx, probeTimeout) defer cancel() ok, err := client.HealthCheck(probeCtx) @@ -220,7 +218,7 @@ func NewReplicaReconciler(opts ReplicaReconcilerOptions) *ReplicaReconciler { } prober := opts.Prober if prober == nil { - prober = grpcModelProber{token: opts.RegistrationToken} + prober = grpcModelProber{clients: &tokenClientFactory{token: opts.RegistrationToken}} } pressureThreshold := opts.PressureThreshold if pressureThreshold == 0 { @@ -352,14 +350,14 @@ func (rc *ReplicaReconciler) drainPendingBackendOps(ctx context.Context) { // Pending-op drain for admin upgrade — fires backend.upgrade so // the slow re-pull doesn't head-of-line-block install traffic on // the same worker. Falls back to the legacy backend.install - // Force=true path on nats.ErrNoResponders for old workers that + // Force=true path on ErrNoRoute for old workers that // don't subscribe to backend.upgrade yet (rolling-update window). // Reconciler retries are background reconciliation with no live // admin watching a progress bar, so opID/onProgress are empty — // the adapter skips the progress subscription entirely. reply, err := rc.adapter.UpgradeBackend(op.NodeID, op.Backend, string(op.Galleries), "", "", "", 0, "", nil) if err != nil { - if errors.Is(err, nats.ErrNoResponders) { + if errors.Is(err, ErrNoRoute) { instReply, instErr := rc.adapter.installWithForceFallback(op.NodeID, op.Backend, string(op.Galleries), "", "", "", 0, "", nil) if instErr != nil { applyErr = instErr @@ -387,14 +385,14 @@ func (rc *ReplicaReconciler) drainPendingBackendOps(ctx context.Context) { continue } - // ErrNoResponders means the node has no active NATS subscription for - // this subject. Either its connection dropped, or it's the wrong - // node type entirely. Mark unhealthy so the health monitor's + // ErrNoRoute means nothing is listening for this subject on the node. + // Either its connection dropped, or it's the wrong node type + // entirely. Mark unhealthy so the health monitor's // heartbeat-only pass doesn't immediately flip it back — and so // ListDuePendingBackendOps (which filters by status=healthy) stops // picking the row until the node genuinely recovers. - if errors.Is(applyErr, nats.ErrNoResponders) { - xlog.Warn("Reconciler: no NATS responders — marking node unhealthy", + if errors.Is(applyErr, ErrNoRoute) { + xlog.Warn("Reconciler: no route to node, marking it unhealthy", "op", op.Op, "backend", op.Backend, "node", op.NodeID) _ = rc.registry.MarkUnhealthy(ctx, op.NodeID) } @@ -480,7 +478,7 @@ func (rc *ReplicaReconciler) probeLoadedModels(ctx context.Context) { return } seen[m.ID] = struct{}{} - switch rc.prober.Probe(ctx, m.Address) { + switch rc.prober.Probe(ctx, m.NodeID, m.Address) { case ProbeAlive: rc.clearProbeFailures(m.ID) // Bump updated_at so we don't probe this row again immediately. @@ -563,7 +561,7 @@ func (rc *ReplicaReconciler) sweepLeakedInFlight(ctx context.Context) { return } seen[m.ID] = struct{}{} - if rc.prober.Probe(ctx, m.Address) != ProbeAlive { + if rc.prober.Probe(ctx, m.NodeID, m.Address) != ProbeAlive { // Busy or unreachable. Busy means the counter may well be real; // unreachable is the reaper's business, not the sweeper's. rc.clearInFlightIdle(m.ID) @@ -725,7 +723,7 @@ type replicaKey struct { } // replyError safely extracts the error text from a possibly-nil reply. -func replyError(reply *messaging.ModelsRunningReply) string { +func replyError(reply *workerctl.ModelsRunningReply) string { if reply == nil { return "nil reply" } diff --git a/core/services/nodes/reconciler_prober_test.go b/core/services/nodes/reconciler_prober_test.go new file mode 100644 index 000000000..7e19bb999 --- /dev/null +++ b/core/services/nodes/reconciler_prober_test.go @@ -0,0 +1,40 @@ +package nodes + +import ( + "context" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + + grpc "github.com/mudler/LocalAI/pkg/grpc" +) + +var _ = Describe("grpcModelProber", func() { + It("builds its client through the factory with the node id and a non-parallel client", func() { + f := &recordingFactory{next: func() grpc.Backend { return &fakeBackendClient{healthy: true} }} + p := grpcModelProber{clients: f} + + out := p.Probe(context.Background(), "n1", "10.0.0.1:50052") + + Expect(f.calls()).To(Equal([]string{"n1@10.0.0.1:50052"})) + Expect(f.parallelFlags()).To(Equal([]bool{false})) + Expect(out).To(Equal(ProbeAlive)) + }) + + DescribeTable("classifies the health answer", + func(mk func() *fakeBackendClient, want ProbeOutcome) { + f := &recordingFactory{next: func() grpc.Backend { return mk() }} + Expect(grpcModelProber{clients: f}.Probe(context.Background(), "n1", "a:1")).To(Equal(want)) + }, + Entry("healthy", func() *fakeBackendClient { return &fakeBackendClient{healthy: true} }, ProbeAlive), + Entry("answered but not healthy", func() *fakeBackendClient { return &fakeBackendClient{healthy: false} }, ProbeUnreachable), + Entry("nothing listening", func() *fakeBackendClient { + return &fakeBackendClient{err: status.Error(codes.Unavailable, "down")} + }, ProbeUnreachable), + Entry("no answer in time", func() *fakeBackendClient { + return &fakeBackendClient{err: context.DeadlineExceeded} + }, ProbeBusy), + ) +}) diff --git a/core/services/nodes/reconciler_test.go b/core/services/nodes/reconciler_test.go index 049fb9441..9246bccb5 100644 --- a/core/services/nodes/reconciler_test.go +++ b/core/services/nodes/reconciler_test.go @@ -734,14 +734,17 @@ var _ = Describe("ReplicaReconciler", func() { }) // fakeProber lets tests control how a model's gRPC address "responds". -// Addresses with no entry default to ProbeUnreachable. +// Addresses with no entry default to ProbeUnreachable. It records the node id +// of every probe so specs can check the reconciler names the node it dials. type fakeProber struct { outcomes map[string]ProbeOutcome calls int + nodeIDs []string } -func (f *fakeProber) Probe(_ context.Context, address string) ProbeOutcome { +func (f *fakeProber) Probe(_ context.Context, nodeID, address string) ProbeOutcome { f.calls++ + f.nodeIDs = append(f.nodeIDs, nodeID) if f.outcomes == nil { return ProbeUnreachable } @@ -838,6 +841,31 @@ var _ = Describe("ReplicaReconciler — state reconciliation", func() { Expect(db.First(&after, "id = ?", "stale-2").Error).To(Succeed()) Expect(after.UpdatedAt).To(BeTemporally("~", time.Now(), time.Second)) }) + + It("probes each replica under the node id of its row", func() { + node := &BackendNode{Name: "n1", NodeType: NodeTypeBackend, Address: "10.0.0.1:50051"} + Expect(registry.Register(context.Background(), node, true)).To(Succeed()) + Expect(db.Create(&NodeModel{ + ID: "stale-3", + NodeID: node.ID, + ModelName: "probed-model", + Address: "10.0.0.1:12345", + State: "loaded", + UpdatedAt: time.Now().Add(-5 * time.Minute), + }).Error).To(Succeed()) + + prober := &fakeProber{outcomes: map[string]ProbeOutcome{"10.0.0.1:12345": ProbeAlive}} + rc := NewReplicaReconciler(ReplicaReconcilerOptions{ + Registry: registry, + DB: db, + Prober: prober, + ProbeStaleAfter: 2 * time.Minute, + }) + + rc.probeLoadedModels(context.Background()) + + Expect(prober.nodeIDs).To(Equal([]string{node.ID})) + }) }) Describe("UpsertPendingBackendOp + RecordPendingBackendOpFailure", func() { diff --git a/core/services/nodes/reconciler_worker_processes_test.go b/core/services/nodes/reconciler_worker_processes_test.go index 8fd848b1d..58d8e72e7 100644 --- a/core/services/nodes/reconciler_worker_processes_test.go +++ b/core/services/nodes/reconciler_worker_processes_test.go @@ -10,23 +10,23 @@ import ( . "github.com/onsi/gomega" "gorm.io/gorm" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" ) // fakeProcessLister stands in for the NATS round-trip to a worker. type fakeProcessLister struct { - running map[string][]messaging.RunningModelInfo + running map[string][]workerctl.RunningModelInfo err error calls int } -func (f *fakeProcessLister) ListRunningModels(nodeID string) (*messaging.ModelsRunningReply, error) { +func (f *fakeProcessLister) ListRunningModels(nodeID string) (*workerctl.ModelsRunningReply, error) { f.calls++ if f.err != nil { return nil, f.err } - return &messaging.ModelsRunningReply{Models: f.running[nodeID]}, nil + return &workerctl.ModelsRunningReply{Models: f.running[nodeID]}, nil } // The worker owns the backend processes, so its answer is authoritative and, @@ -80,7 +80,7 @@ var _ = Describe("ReplicaReconciler — reconcile against worker processes", fun It("reaps a row for a model the worker is not running", func() { seed("ghost-1", "ghost-model", 0, 5*time.Minute) - lister := &fakeProcessLister{running: map[string][]messaging.RunningModelInfo{}} + lister := &fakeProcessLister{running: map[string][]workerctl.RunningModelInfo{}} rc := newReconciler(lister) for range workerMissesBeforeReap { @@ -95,7 +95,7 @@ var _ = Describe("ReplicaReconciler — reconcile against worker processes", fun It("keeps and refreshes a row the worker confirms is running", func() { seed("live-1", "live-model", 0, 5*time.Minute) - lister := &fakeProcessLister{running: map[string][]messaging.RunningModelInfo{ + lister := &fakeProcessLister{running: map[string][]workerctl.RunningModelInfo{ node.ID: {{ModelID: "live-model", ReplicaIndex: 0, Address: "10.0.0.1:12345"}}, }} rc := newReconciler(lister) @@ -113,7 +113,7 @@ var _ = Describe("ReplicaReconciler — reconcile against worker processes", fun // This is what keeps a busy backend off the port prober entirely: the // worker vouches for it, so it never looks stale enough to probe. seed("busy-1", "busy-model", 0, 5*time.Minute) - lister := &fakeProcessLister{running: map[string][]messaging.RunningModelInfo{ + lister := &fakeProcessLister{running: map[string][]workerctl.RunningModelInfo{ node.ID: {{ModelID: "busy-model", ReplicaIndex: 0, Address: "10.0.0.1:12345"}}, }} rc := newReconciler(lister) @@ -128,7 +128,7 @@ var _ = Describe("ReplicaReconciler — reconcile against worker processes", fun It("distinguishes replicas of the same model", func() { seed("rep-0", "multi-model", 0, 5*time.Minute) seed("rep-1", "multi-model", 1, 5*time.Minute) - lister := &fakeProcessLister{running: map[string][]messaging.RunningModelInfo{ + lister := &fakeProcessLister{running: map[string][]workerctl.RunningModelInfo{ node.ID: {{ModelID: "multi-model", ReplicaIndex: 0, Address: "10.0.0.1:12345"}}, }} rc := newReconciler(lister) @@ -168,7 +168,7 @@ var _ = Describe("ReplicaReconciler — reconcile against worker processes", fun // A row created moments ago may legitimately not be in the worker's // table yet; judging it immediately would race every fresh load. seed("fresh-1", "fresh-model", 0, 0) - lister := &fakeProcessLister{running: map[string][]messaging.RunningModelInfo{}} + lister := &fakeProcessLister{running: map[string][]workerctl.RunningModelInfo{}} rc := newReconciler(lister) for range workerMissesBeforeReap + 2 { @@ -182,7 +182,7 @@ var _ = Describe("ReplicaReconciler — reconcile against worker processes", fun It("requires consecutive misses before reaping", func() { seed("flap-1", "flap-model", 0, 5*time.Minute) - lister := &fakeProcessLister{running: map[string][]messaging.RunningModelInfo{}} + lister := &fakeProcessLister{running: map[string][]workerctl.RunningModelInfo{}} rc := newReconciler(lister) for i := 1; i < workerMissesBeforeReap; i++ { @@ -193,7 +193,7 @@ var _ = Describe("ReplicaReconciler — reconcile against worker processes", fun } // It shows up again: the streak resets. - lister.running[node.ID] = []messaging.RunningModelInfo{ + lister.running[node.ID] = []workerctl.RunningModelInfo{ {ModelID: "flap-model", ReplicaIndex: 0, Address: "10.0.0.1:12345"}, } rc.reconcileNodeProcesses(context.Background()) diff --git a/core/services/nodes/registry.go b/core/services/nodes/registry.go index ac96e4d9d..f6daf36cf 100644 --- a/core/services/nodes/registry.go +++ b/core/services/nodes/registry.go @@ -1298,7 +1298,7 @@ func (r *NodeRegistry) GetByName(ctx context.Context, name string) (*BackendNode } // MarkUnhealthy sets a node status to unhealthy. Deliberately status-only: -// callers fire this on transient triggers (a single nats.ErrNoResponders from +// callers fire this on transient triggers (a single ErrNoRoute from // managers_distributed / reconciler) where the next heartbeat is expected to // flip the node back to healthy, and cascade-deleting node_models here would // force a full model reload on every brief NATS hiccup. Stale rows are reaped @@ -2767,8 +2767,8 @@ func (r *NodeRegistry) DeleteStalePendingBackendOps(ctx context.Context, grace t cutoff := time.Now().Add(-grace) // Draining nodes are cleared immediately (admin action; model rows already // purged). Offline AND unhealthy nodes are cleared only once their heartbeat - // is older than the grace window: a node marked unhealthy on a NATS - // ErrNoResponders never transitions to offline (health.go skips re-marking + // is older than the grace window: a node marked unhealthy on an + // ErrNoRoute never transitions to offline (health.go skips re-marking // it), so without including unhealthy here its ops would leak exactly like // the offline case. A node with a fresh heartbeat (last_heartbeat > cutoff) // is recovering and keeps its op for retry. diff --git a/core/services/nodes/revision_eligibility_test.go b/core/services/nodes/revision_eligibility_test.go index 01f3bfd37..5d3e5f7f1 100644 --- a/core/services/nodes/revision_eligibility_test.go +++ b/core/services/nodes/revision_eligibility_test.go @@ -9,8 +9,8 @@ import ( . "github.com/onsi/gomega" "gorm.io/gorm" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" ) @@ -183,15 +183,17 @@ var _ = Describe("revision eligibility consumers", func() { rc := NewReplicaReconciler(ReplicaReconcilerOptions{Registry: registry, DB: db, Prober: prober, ProbeStaleAfter: time.Minute}) rc.probeLoadedModels(ctx) Expect(prober.addresses).To(ConsistOf("current")) + Expect(prober.nodeIDs).To(ConsistOf(nodes["current"].ID)) }), Entry("sweepLeakedInFlight", func(prober *recordingEligibilityProber, _ *recordingEligibilityLister) { Expect(db.Model(&NodeModel{}).Where("model_name = ?", modelName).Updates(map[string]any{"in_flight": 1, "last_used": time.Now().Add(-2 * inFlightLeakIdleAfter)}).Error).To(Succeed()) rc := NewReplicaReconciler(ReplicaReconcilerOptions{Registry: registry, DB: db, Prober: prober}) rc.sweepLeakedInFlight(ctx) Expect(prober.addresses).To(ConsistOf("current")) + Expect(prober.nodeIDs).To(ConsistOf(nodes["current"].ID)) }), Entry("reconcileNodeProcesses", func(_ *recordingEligibilityProber, lister *recordingEligibilityLister) { - lister.running = map[string][]messaging.RunningModelInfo{nodes["current"].ID: {{ModelID: modelName, ReplicaIndex: 0}}} + lister.running = map[string][]workerctl.RunningModelInfo{nodes["current"].ID: {{ModelID: modelName, ReplicaIndex: 0}}} rc := NewReplicaReconciler(ReplicaReconcilerOptions{Registry: registry, DB: db, ProcessLister: lister, ProbeStaleAfter: time.Minute}) rc.reconcileNodeProcesses(ctx) Expect(lister.nodeIDs).To(ConsistOf(nodes["current"].ID)) @@ -262,19 +264,20 @@ var _ = Describe("revision eligibility consumers", func() { }) }) -type recordingEligibilityProber struct{ addresses []string } +type recordingEligibilityProber struct{ addresses, nodeIDs []string } -func (p *recordingEligibilityProber) Probe(_ context.Context, address string) ProbeOutcome { +func (p *recordingEligibilityProber) Probe(_ context.Context, nodeID, address string) ProbeOutcome { p.addresses = append(p.addresses, address) + p.nodeIDs = append(p.nodeIDs, nodeID) return ProbeAlive } type recordingEligibilityLister struct { nodeIDs []string - running map[string][]messaging.RunningModelInfo + running map[string][]workerctl.RunningModelInfo } -func (l *recordingEligibilityLister) ListRunningModels(nodeID string) (*messaging.ModelsRunningReply, error) { +func (l *recordingEligibilityLister) ListRunningModels(nodeID string) (*workerctl.ModelsRunningReply, error) { l.nodeIDs = append(l.nodeIDs, nodeID) - return &messaging.ModelsRunningReply{Models: l.running[nodeID]}, nil + return &workerctl.ModelsRunningReply{Models: l.running[nodeID]}, nil } diff --git a/core/services/nodes/router.go b/core/services/nodes/router.go index 84cad29a6..e858fb7b6 100644 --- a/core/services/nodes/router.go +++ b/core/services/nodes/router.go @@ -1394,7 +1394,7 @@ func (r *SmartRouter) installBackendOnNode(ctx context.Context, node *BackendNod } func (r *SmartRouter) buildClientForAddr(node *BackendNode, addr string, parallel bool) grpc.Backend { - client := r.clientFactory.NewClient(addr, parallel) + client := r.clientFactory.NewClient(node.ID, addr, parallel) // Wrap with file staging if configured if r.fileStager != nil { diff --git a/core/services/nodes/router_liveness.go b/core/services/nodes/router_liveness.go index 88646162f..fd0836657 100644 --- a/core/services/nodes/router_liveness.go +++ b/core/services/nodes/router_liveness.go @@ -5,7 +5,6 @@ import ( "errors" "github.com/mudler/xlog" - "github.com/nats-io/nats.go" ) // maxNodeLivenessRetries bounds how many unreachable nodes a single scheduling @@ -16,7 +15,7 @@ const maxNodeLivenessRetries = 3 // nodeAnswersOnBus reports whether a node still has a live subscription. // -// Only nats.ErrNoResponders means "absent". Any other outcome, a timeout or a +// Only ErrNoRoute means "absent". Any other outcome, a timeout or a // transport hiccup, leaves the node eligible: wrongly excluding a node that is // merely slow costs real capacity, while the install that follows already // reports its own failure. When no command sender is configured there is no bus @@ -27,7 +26,7 @@ func (r *SmartRouter) nodeAnswersOnBus(node *BackendNode) bool { return true } err := r.unloader.PingNode(node.ID) - return !errors.Is(err, nats.ErrNoResponders) + return !errors.Is(err, ErrNoRoute) } // pickReachableNode calls selectNode until it yields a node that still answers diff --git a/core/services/nodes/router_liveness_route_test.go b/core/services/nodes/router_liveness_route_test.go new file mode 100644 index 000000000..5bfb782de --- /dev/null +++ b/core/services/nodes/router_liveness_route_test.go @@ -0,0 +1,44 @@ +package nodes + +import ( + "context" + "errors" + "fmt" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +// pingStub is a NodeCommandSender that answers PingNode and nothing else. +type pingStub struct { + NodeCommandSender + err error +} + +func (p pingStub) PingNode(string) error { return p.err } + +var _ = Describe("Scheduler liveness on the control path", func() { + node := &BackendNode{ID: "n1", Name: "n1"} + + answers := func(err error) bool { + r := &SmartRouter{unloader: pingStub{err: err}} + return r.nodeAnswersOnBus(node) + } + + It("excludes a node only when there is no route to it", func() { + Expect(answers(fmt.Errorf("wrapped: %w", ErrNoRoute))).To(BeFalse()) + }) + + It("keeps a node whose ping timed out", func() { + Expect(answers(context.DeadlineExceeded)).To(BeTrue()) + }) + + It("keeps a node whose ping failed for any other reason", func() { + Expect(answers(errors.New("connection reset"))).To(BeTrue()) + }) + + It("keeps every node when no command sender is configured", func() { + r := &SmartRouter{} + Expect(r.nodeAnswersOnBus(node)).To(BeTrue()) + }) +}) diff --git a/core/services/nodes/router_load_budget_test.go b/core/services/nodes/router_load_budget_test.go index 921b19905..744cddd0e 100644 --- a/core/services/nodes/router_load_budget_test.go +++ b/core/services/nodes/router_load_budget_test.go @@ -11,7 +11,7 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" grpc "github.com/mudler/LocalAI/pkg/grpc" pb "github.com/mudler/LocalAI/pkg/grpc/proto" ggrpc "google.golang.org/grpc" @@ -84,7 +84,7 @@ func (b *holdBackend) ctxErrAtEnd() error { type holdClientFactory struct{ client *holdBackend } -func (f *holdClientFactory) NewClient(_ string, _ bool) grpc.Backend { return f.client } +func (f *holdClientFactory) NewClient(_, _ string, _ bool) grpc.Backend { return f.client } var _ = Describe("size-derived remote LoadModel budget", func() { // Production, on an NVIDIA Jetson Thor worker: a 70 GB video checkpoint @@ -108,7 +108,7 @@ var _ = Describe("size-derived remote LoadModel budget", func() { backend = &holdBackend{} factory = &holdClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, } dir = GinkgoT().TempDir() }) diff --git a/core/services/nodes/router_load_job_test.go b/core/services/nodes/router_load_job_test.go index 65d1ef939..e5eea25ca 100644 --- a/core/services/nodes/router_load_job_test.go +++ b/core/services/nodes/router_load_job_test.go @@ -10,8 +10,8 @@ import ( . "github.com/onsi/gomega" "github.com/mudler/LocalAI/core/config" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "gorm.io/gorm" ) @@ -51,7 +51,7 @@ var _ = Describe("Route cold-load jobs", func() { backend = &stubBackend{healthResult: true, loadResult: &pb.Result{Success: true}} factory = &stubClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, } }) @@ -123,7 +123,7 @@ var _ = Describe("Route cold-load jobs", func() { It("reports the load's real failure to every waiter", func() { release := make(chan struct{}) unloader.installHook = func() { <-release } - unloader.installReply = &messaging.BackendInstallReply{Success: false, Error: "worker out of disk"} + unloader.installReply = &workerctl.BackendInstallReply{Success: false, Error: "worker out of disk"} router := newRouter() diff --git a/core/services/nodes/router_load_timeout_test.go b/core/services/nodes/router_load_timeout_test.go index 295dc35d8..7c8bdfa77 100644 --- a/core/services/nodes/router_load_timeout_test.go +++ b/core/services/nodes/router_load_timeout_test.go @@ -10,7 +10,7 @@ import ( . "github.com/onsi/gomega" "github.com/mudler/LocalAI/core/config" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" grpc "github.com/mudler/LocalAI/pkg/grpc" pb "github.com/mudler/LocalAI/pkg/grpc/proto" ggrpc "google.golang.org/grpc" @@ -51,7 +51,7 @@ func (b *deadlineBackend) budget() time.Duration { type deadlineClientFactory struct{ client *deadlineBackend } -func (f *deadlineClientFactory) NewClient(_ string, _ bool) grpc.Backend { return f.client } +func (f *deadlineClientFactory) NewClient(_, _ string, _ bool) grpc.Backend { return f.client } var _ = Describe("remote LoadModel deadline", func() { var ( @@ -67,7 +67,7 @@ var _ = Describe("remote LoadModel deadline", func() { backend = &deadlineBackend{} factory = &deadlineClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, } }) diff --git a/core/services/nodes/router_reap_load_test.go b/core/services/nodes/router_reap_load_test.go index 67376c06f..ee5ce0fa9 100644 --- a/core/services/nodes/router_reap_load_test.go +++ b/core/services/nodes/router_reap_load_test.go @@ -9,7 +9,7 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" grpc "github.com/mudler/LocalAI/pkg/grpc" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "google.golang.org/grpc/codes" @@ -43,7 +43,7 @@ func (b *failingLoadBackend) LoadModel(_ context.Context, _ *pb.ModelOptions, _ type failingClientFactory struct{ client *failingLoadBackend } -func (f *failingClientFactory) NewClient(_ string, _ bool) grpc.Backend { return f.client } +func (f *failingClientFactory) NewClient(_, _ string, _ bool) grpc.Backend { return f.client } // replicaSlotRouter pins the replica slot scheduleAndLoad allocates so a spec // can assert the reaped process key carries the real index, not a hardcoded 0. @@ -69,7 +69,7 @@ var _ = Describe("reaping an abandoned remote load", func() { reg = &replicaSlotRouter{fakeModelRouter: base, replica: 2} backend = &failingLoadBackend{} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, } }) diff --git a/core/services/nodes/router_revision_lifecycle_test.go b/core/services/nodes/router_revision_lifecycle_test.go index 7b1de6144..5adefe92d 100644 --- a/core/services/nodes/router_revision_lifecycle_test.go +++ b/core/services/nodes/router_revision_lifecycle_test.go @@ -9,8 +9,8 @@ import ( corebackend "github.com/mudler/LocalAI/core/backend" "github.com/mudler/LocalAI/core/config" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/testutil" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/system" @@ -27,11 +27,11 @@ type recordingRevisionStopper struct { err error } -func (s *recordingRevisionStopper) StopModelReplica(_ context.Context, _ string, replica NodeModel, _ bool) (messaging.ModelStopReply, error) { +func (s *recordingRevisionStopper) StopModelReplica(_ context.Context, _ string, replica NodeModel, _ bool) (workerctl.ModelStopReply, error) { s.mu.Lock() s.replicas = append(s.replicas, replica) s.mu.Unlock() - return messaging.ModelStopReply{}, s.err + return workerctl.ModelStopReply{}, s.err } var _ = Describe("revision-bound load publication", func() { @@ -56,7 +56,7 @@ var _ = Describe("revision-bound load publication", func() { node = &BackendNode{Name: "revision-worker", 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: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}} + unloader = &fakeUnloader{installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}} }) It("quarantines and exactly stops a load that finishes after its revision changes", func() { diff --git a/core/services/nodes/router_slot_uncertainty_test.go b/core/services/nodes/router_slot_uncertainty_test.go index b32649197..09e4110e5 100644 --- a/core/services/nodes/router_slot_uncertainty_test.go +++ b/core/services/nodes/router_slot_uncertainty_test.go @@ -7,7 +7,7 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" "gorm.io/gorm" ) @@ -37,7 +37,7 @@ var _ = Describe("replica slot lookup under database latency", func() { backend = &stubBackend{loadResult: &pb.Result{Success: true}} factory = &stubClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.10:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.10:9001"}, } }) @@ -90,7 +90,7 @@ var _ = Describe("node selection under database latency", func() { reg = &fakeModelRouter{findAndLockErr: errors.New("not found")} factory = &stubClientFactory{client: &stubBackend{loadResult: &pb.Result{Success: true}}} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.10:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.10:9001"}, } }) diff --git a/core/services/nodes/router_staging_context_test.go b/core/services/nodes/router_staging_context_test.go index f0b07a7a5..1639663ab 100644 --- a/core/services/nodes/router_staging_context_test.go +++ b/core/services/nodes/router_staging_context_test.go @@ -9,7 +9,7 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" ) @@ -51,7 +51,7 @@ var _ = Describe("Route cold-load staging context", func() { } backend := &stubBackend{loadResult: &pb.Result{Success: true}} factory := &stubClientFactory{client: backend} - unloader := &fakeUnloader{installReply: &messaging.BackendInstallReply{ + unloader := &fakeUnloader{installReply: &workerctl.BackendInstallReply{ Success: true, Address: "10.0.0.1:9001", }} diff --git a/core/services/nodes/router_staging_deadline_test.go b/core/services/nodes/router_staging_deadline_test.go index 35d1d2ae5..316a618b5 100644 --- a/core/services/nodes/router_staging_deadline_test.go +++ b/core/services/nodes/router_staging_deadline_test.go @@ -11,7 +11,7 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" ) @@ -96,7 +96,7 @@ var _ = Describe("cold-load staging deadline", func() { findIdleNode: &BackendNode{ID: "n1", Name: "nvidia-thor", Address: "10.0.0.1:50051"}, } factory = &stubClientFactory{client: &stubBackend{loadResult: &pb.Result{Success: true}}} - unloader = &fakeUnloader{installReply: &messaging.BackendInstallReply{ + unloader = &fakeUnloader{installReply: &workerctl.BackendInstallReply{ Success: true, Address: "10.0.0.1:9001", }} diff --git a/core/services/nodes/router_test.go b/core/services/nodes/router_test.go index 7577e2846..a52d36bf9 100644 --- a/core/services/nodes/router_test.go +++ b/core/services/nodes/router_test.go @@ -12,13 +12,12 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes/prefixcache" "github.com/mudler/LocalAI/core/services/testutil" + "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" - "github.com/nats-io/nats.go" ggrpc "google.golang.org/grpc" "google.golang.org/protobuf/proto" "gorm.io/gorm" @@ -476,7 +475,7 @@ type stubClientFactory struct { client *stubBackend } -func (f *stubClientFactory) NewClient(_ string, _ bool) grpc.Backend { +func (f *stubClientFactory) NewClient(_, _ string, _ bool) grpc.Backend { return f.client } @@ -489,7 +488,7 @@ type fakeUnloader struct { // goroutines (e.g. singleflight specs) don't race the slice appends. mu sync.Mutex - installReply *messaging.BackendInstallReply + installReply *workerctl.BackendInstallReply installErr error installCalls []installCall // every InstallBackend invocation, in order // installHook, if non-nil, runs at the start of InstallBackend before @@ -498,7 +497,7 @@ type fakeUnloader struct { // blocks on a channel to overlap two callers. installHook func() - upgradeReply *messaging.BackendUpgradeReply + upgradeReply *workerctl.BackendUpgradeReply upgradeErr error upgradeCalls []upgradeCall // every UpgradeBackend invocation, in order @@ -533,7 +532,7 @@ type upgradeCall struct { replica int } -func (f *fakeUnloader) InstallBackend(nodeID, backend, modelID, _, _, _, _ string, replica int, _ string, _ func(messaging.BackendInstallProgressEvent)) (*messaging.BackendInstallReply, error) { +func (f *fakeUnloader) InstallBackend(nodeID, backend, modelID, _, _, _, _ string, replica int, _ string, _ func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { // installHook intentionally runs OUTSIDE the mutex: the hook may block // on a channel and we don't want to serialize concurrent callers, // which would defeat the singleflight-overlap test. @@ -546,19 +545,19 @@ func (f *fakeUnloader) InstallBackend(nodeID, backend, modelID, _, _, _, _ strin return f.installReply, f.installErr } -func (f *fakeUnloader) UpgradeBackend(nodeID, backend, _, _, _, _ string, replica int, _ string, _ func(messaging.BackendInstallProgressEvent)) (*messaging.BackendUpgradeReply, error) { +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}) f.mu.Unlock() return f.upgradeReply, f.upgradeErr } -func (f *fakeUnloader) DeleteBackend(_, _ string) (*messaging.BackendDeleteReply, error) { - return &messaging.BackendDeleteReply{Success: true}, nil +func (f *fakeUnloader) DeleteBackend(_, _ string) (*workerctl.BackendDeleteReply, error) { + return &workerctl.BackendDeleteReply{Success: true}, nil } -func (f *fakeUnloader) ListBackends(_ string) (*messaging.BackendListReply, error) { - return &messaging.BackendListReply{}, nil +func (f *fakeUnloader) ListBackends(_ string) (*workerctl.BackendListReply, error) { + return &workerctl.BackendListReply{}, nil } func (f *fakeUnloader) StopBackend(nodeID, backend string) error { @@ -582,7 +581,7 @@ func (f *fakeUnloader) PingNode(nodeID string) error { dead := f.deadNodes[nodeID] f.mu.Unlock() if dead { - return nats.ErrNoResponders + return ErrNoRoute } return f.pingErr } @@ -608,7 +607,7 @@ var _ = Describe("SmartRouter", func() { backend = &stubBackend{} factory = &stubClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{ + installReply: &workerctl.BackendInstallReply{ Success: true, Address: "10.0.0.1:9001", }, @@ -759,7 +758,7 @@ var _ = Describe("SmartRouter", func() { } factory = &stubClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{ + installReply: &workerctl.BackendInstallReply{ Success: true, Address: "10.0.0.1:9001", }, @@ -941,7 +940,7 @@ var _ = Describe("SmartRouter", func() { } factory = &stubClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{ + installReply: &workerctl.BackendInstallReply{ Success: true, Address: "10.0.0.1:9001", }, @@ -1040,7 +1039,7 @@ var _ = Describe("SmartRouter", func() { } factory := &stubClientFactory{client: backend} unloader := &fakeUnloader{ - installReply: &messaging.BackendInstallReply{ + installReply: &workerctl.BackendInstallReply{ Success: true, Address: "10.0.0.71:9001", }, @@ -1310,7 +1309,7 @@ var _ = Describe("SmartRouter", func() { started := make(chan struct{}, 5) release := make(chan struct{}) unloader := &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}, } unloader.installHook = func() { started <- struct{}{} @@ -1349,7 +1348,7 @@ var _ = Describe("SmartRouter", func() { It("does NOT coalesce installs for different (modelID, replica) keys", func() { node := &BackendNode{ID: "n1", Name: "node-1", Address: "10.0.0.1:50051"} unloader := &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:50100"}, } router := NewSmartRouter(&fakeModelRouter{}, SmartRouterOptions{ Unloader: unloader, @@ -1423,7 +1422,7 @@ var _ = Describe("SmartRouter prefix-cache routing", func() { backend = &stubBackend{healthResult: true} factory = &stubClientFactory{client: backend} unloader = &fakeUnloader{ - installReply: &messaging.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, + installReply: &workerctl.BackendInstallReply{Success: true, Address: "10.0.0.1:9001"}, } }) diff --git a/core/services/nodes/staging_progress.go b/core/services/nodes/staging_progress.go index 0a6ddc50e..c6bfafedb 100644 --- a/core/services/nodes/staging_progress.go +++ b/core/services/nodes/staging_progress.go @@ -87,7 +87,7 @@ func (t *StagingTracker) SetPublisher(p messaging.Publisher) { // SubscribeBroadcasts subscribes to peer replicas' staging-progress broadcasts // and mirrors them into this tracker, so /api/operations on any replica surfaces // staging ops it did not originate. Returns the subscription for cleanup. -func (t *StagingTracker) SubscribeBroadcasts(nc messaging.MessagingClient) (messaging.Subscription, error) { +func (t *StagingTracker) SubscribeBroadcasts(nc messaging.Broadcaster) (messaging.Subscription, error) { return messaging.SubscribeJSON(nc, messaging.SubjectStagingProgressWildcard, func(evt StagingProgressEvent) { if evt.ModelID == "" { return diff --git a/core/services/nodes/unloader.go b/core/services/nodes/unloader.go index b95b1330b..ba6bc08b6 100644 --- a/core/services/nodes/unloader.go +++ b/core/services/nodes/unloader.go @@ -5,13 +5,11 @@ import ( "encoding/json" "errors" "fmt" - "strings" "time" - "github.com/nats-io/nats.go" - "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/xlog" ) @@ -26,19 +24,21 @@ import ( // UpgradeBackend is the destructive force-reinstall path: the worker stops // every live process for the backend, re-pulls the gallery artifact, and // replies. Caller (DistributedBackendManager.UpgradeBackend) handles -// rolling-update fallback to the legacy install Force=true path on -// nats.ErrNoResponders for old workers that don't subscribe to the new -// backend.upgrade subject. +// rolling-update fallback to the legacy install Force=true path. +// +// PingNode returns ErrNoRoute when nothing answers for the node, which is the +// only condition callers may read as "this node cannot be given work". +// UpgradeBackend returns ErrNoRoute on an old worker that does not serve +// backend.upgrade, and the caller falls back to the legacy install. type NodeCommandSender interface { - InstallBackend(nodeID, backendType, modelID, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(messaging.BackendInstallProgressEvent)) (*messaging.BackendInstallReply, error) - UpgradeBackend(nodeID, backendType, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(messaging.BackendInstallProgressEvent)) (*messaging.BackendUpgradeReply, error) - DeleteBackend(nodeID, backendName string) (*messaging.BackendDeleteReply, error) - ListBackends(nodeID string) (*messaging.BackendListReply, error) + InstallBackend(nodeID, backendType, modelID, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) + UpgradeBackend(nodeID, backendType, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendUpgradeReply, error) + DeleteBackend(nodeID, backendName string) (*workerctl.BackendDeleteReply, error) + ListBackends(nodeID string) (*workerctl.BackendListReply, error) StopBackend(nodeID, backend string) error UnloadModelOnNode(nodeID, modelName string) error // PingNode reports whether the node is still subscribed on the bus. It - // returns nats.ErrNoResponders when nothing answers for the node, which is - // the only condition callers may read as "this node cannot be given work". + // returns ErrNoRoute when nothing answers for the node. PingNode(nodeID string) error } @@ -93,7 +93,7 @@ const exactModelStopTimeout = 10 * time.Second // StopModelReplica stops only the process represented by replica. Configuration // 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) (messaging.ModelStopReply, error) { +func (a *RemoteUnloaderAdapter) StopModelReplica(ctx context.Context, nodeID string, replica NodeModel, force bool) (workerctl.ModelStopReply, error) { if ctx == nil { ctx = context.Background() } @@ -101,12 +101,12 @@ func (a *RemoteUnloaderAdapter) StopModelReplica(ctx context.Context, nodeID str defer cancel() type result struct { - reply *messaging.ModelStopReply + reply *workerctl.ModelStopReply err error } done := make(chan result, 1) go func() { - reply, err := messaging.RequestJSON[messaging.ModelStopRequest, messaging.ModelStopReply](a.nats, messaging.SubjectNodeModelStop(nodeID), messaging.ModelStopRequest{ + 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, @@ -118,10 +118,10 @@ func (a *RemoteUnloaderAdapter) StopModelReplica(ctx context.Context, nodeID str select { case <-ctx.Done(): - return messaging.ModelStopReply{}, ctx.Err() + return workerctl.ModelStopReply{}, ctx.Err() case result := <-done: if result.err != nil { - return messaging.ModelStopReply{}, result.err + return workerctl.ModelStopReply{}, result.err } return *result.reply, nil } @@ -213,8 +213,8 @@ func (a *RemoteUnloaderAdapter) InstallBackend( nodeID, backendType, modelID, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, - onProgress func(messaging.BackendInstallProgressEvent), -) (*messaging.BackendInstallReply, error) { + 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) @@ -222,7 +222,7 @@ func (a *RemoteUnloaderAdapter) InstallBackend( // request so we don't miss early events. sub := a.subscribeProgress(nodeID, opID, onProgress) - reply, err := messaging.RequestJSON[messaging.BackendInstallRequest, messaging.BackendInstallReply](a.nats, subject, messaging.BackendInstallRequest{ + reply, err := controlRequestJSON[workerctl.BackendInstallRequest, workerctl.BackendInstallReply](a.nats, subject, workerctl.BackendInstallRequest{ Backend: backendType, ModelID: modelID, BackendGalleries: galleriesJSON, @@ -255,13 +255,13 @@ func (a *RemoteUnloaderAdapter) InstallBackend( // install-progress subject rather than minting a new one (no new NATS // permission, no new rolling-update compat surface). Caller must Unsubscribe // the returned subscription after the request completes. -func (a *RemoteUnloaderAdapter) subscribeProgress(nodeID, opID string, onProgress func(messaging.BackendInstallProgressEvent)) messaging.Subscription { +func (a *RemoteUnloaderAdapter) subscribeProgress(nodeID, opID string, onProgress func(workerctl.BackendInstallProgressEvent)) messaging.Subscription { if onProgress == nil || opID == "" { return nil } progressSubject := messaging.SubjectNodeBackendInstallProgress(nodeID, opID) s, subErr := a.nats.Subscribe(progressSubject, func(raw []byte) { - var ev messaging.BackendInstallProgressEvent + var ev workerctl.BackendInstallProgressEvent if err := json.Unmarshal(raw, &ev); err != nil { xlog.Debug("malformed backend progress event", "subject", progressSubject, "error", err) return @@ -293,13 +293,13 @@ func (a *RemoteUnloaderAdapter) subscribeProgress(nodeID, opID string, onProgres // Timeout: configured via DistributedConfig.BackendUpgradeTimeoutOrDefault // (default 15m). Real-world worst case observed: 8-10 minutes for large // CUDA-l4t backend images on Jetson over WiFi. -func (a *RemoteUnloaderAdapter) UpgradeBackend(nodeID, backendType, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(messaging.BackendInstallProgressEvent)) (*messaging.BackendUpgradeReply, error) { +func (a *RemoteUnloaderAdapter) UpgradeBackend(nodeID, backendType, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendUpgradeReply, error) { subject := messaging.SubjectNodeBackendUpgrade(nodeID) xlog.Info("Sending NATS backend.upgrade", "nodeID", nodeID, "backend", backendType, "replica", replicaIndex, "opID", opID) sub := a.subscribeProgress(nodeID, opID, onProgress) - reply, err := messaging.RequestJSON[messaging.BackendUpgradeRequest, messaging.BackendUpgradeReply](a.nats, subject, messaging.BackendUpgradeRequest{ + reply, err := controlRequestJSON[workerctl.BackendUpgradeRequest, workerctl.BackendUpgradeReply](a.nats, subject, workerctl.BackendUpgradeRequest{ Backend: backendType, BackendGalleries: galleriesJSON, URI: uri, @@ -327,17 +327,17 @@ func (a *RemoteUnloaderAdapter) UpgradeBackend(nodeID, backendType, galleriesJSO // installWithForceFallback is the rolling-update fallback used by // DistributedBackendManager.UpgradeBackend when backend.upgrade returns -// nats.ErrNoResponders (the worker is on a pre-2026-05-08 build that +// ErrNoRoute (the worker is on a pre-2026-05-08 build that // doesn't subscribe to the new subject). It re-fires the legacy // backend.install with Force=true. Drop this once every worker is on // 2026-05-08 or newer. -func (a *RemoteUnloaderAdapter) installWithForceFallback(nodeID, backendType, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(messaging.BackendInstallProgressEvent)) (*messaging.BackendInstallReply, error) { +func (a *RemoteUnloaderAdapter) installWithForceFallback(nodeID, backendType, galleriesJSON, uri, name, alias string, replicaIndex int, opID string, onProgress func(workerctl.BackendInstallProgressEvent)) (*workerctl.BackendInstallReply, error) { subject := messaging.SubjectNodeBackendInstall(nodeID) xlog.Warn("Falling back to legacy backend.install Force=true (old worker)", "nodeID", nodeID, "backend", backendType) sub := a.subscribeProgress(nodeID, opID, onProgress) - reply, err := messaging.RequestJSON[messaging.BackendInstallRequest, messaging.BackendInstallReply](a.nats, subject, messaging.BackendInstallRequest{ + reply, err := controlRequestJSON[workerctl.BackendInstallRequest, workerctl.BackendInstallReply](a.nats, subject, workerctl.BackendInstallRequest{ Backend: backendType, BackendGalleries: galleriesJSON, URI: uri, @@ -362,11 +362,11 @@ func (a *RemoteUnloaderAdapter) installWithForceFallback(nodeID, backendType, ga } // ListBackends queries a worker node for its installed backends via NATS request-reply. -func (a *RemoteUnloaderAdapter) ListBackends(nodeID string) (*messaging.BackendListReply, error) { +func (a *RemoteUnloaderAdapter) ListBackends(nodeID string) (*workerctl.BackendListReply, error) { subject := messaging.SubjectNodeBackendList(nodeID) xlog.Debug("Sending NATS backend.list", "nodeID", nodeID) - return messaging.RequestJSON[messaging.BackendListRequest, messaging.BackendListReply](a.nats, subject, messaging.BackendListRequest{}, 30*time.Second) + return controlRequestJSON[workerctl.BackendListRequest, workerctl.BackendListReply](a.nats, subject, workerctl.BackendListRequest{}, 30*time.Second) } // PingNode checks that a worker still has a live subscription on the bus. @@ -385,7 +385,7 @@ func (a *RemoteUnloaderAdapter) ListBackends(nodeID string) (*messaging.BackendL // it is the safer question to ask. // // A worker that answers anything is alive. Only when every subject reports no -// responders is the node treated as absent, so adding a newer subject here can +// route is the node treated as absent, so adding a newer subject here can // never condemn an older worker. func (a *RemoteUnloaderAdapter) PingNode(nodeID string) error { subjects := []string{ @@ -394,12 +394,12 @@ func (a *RemoteUnloaderAdapter) PingNode(nodeID string) error { } var lastErr error for _, subject := range subjects { - _, err := messaging.RequestJSON[messaging.BackendListRequest, messaging.BackendListReply]( - a.nats, subject, messaging.BackendListRequest{}, 5*time.Second) + _, err := controlRequestJSON[workerctl.BackendListRequest, workerctl.BackendListReply]( + a.nats, subject, workerctl.BackendListRequest{}, 5*time.Second) if err == nil { return nil } - if !errors.Is(err, nats.ErrNoResponders) { + if !errors.Is(err, ErrNoRoute) { // Reached someone, or failed for a reason that is not absence. // Either way the node is not proven gone. return nil @@ -416,10 +416,10 @@ func (a *RemoteUnloaderAdapter) PingNode(nodeID string) error { // in-memory process table, so a slow reply means the worker itself is in // trouble, and the caller treats no-answer as "don't know" rather than as // "nothing running". -func (a *RemoteUnloaderAdapter) ListRunningModels(nodeID string) (*messaging.ModelsRunningReply, error) { +func (a *RemoteUnloaderAdapter) ListRunningModels(nodeID string) (*workerctl.ModelsRunningReply, error) { subject := messaging.SubjectNodeModelsRunning(nodeID) - return messaging.RequestJSON[messaging.ModelsRunningRequest, messaging.ModelsRunningReply]( - a.nats, subject, messaging.ModelsRunningRequest{}, 10*time.Second) + return controlRequestJSON[workerctl.ModelsRunningRequest, workerctl.ModelsRunningReply]( + a.nats, subject, workerctl.ModelsRunningRequest{}, 10*time.Second) } // backendStopAckTimeout bounds the wait for a worker's backend.stop reply. @@ -459,12 +459,12 @@ func (a *RemoteUnloaderAdapter) StopBackend(nodeID, backend string) error { // on that to skip the registry cleanup for a node they could not reach. func (a *RemoteUnloaderAdapter) stopBackend(nodeID, backend string, force bool) error { subject := messaging.SubjectNodeBackendStop(nodeID) - req := messaging.BackendStopRequest{Backend: backend, Force: force} + req := workerctl.BackendStopRequest{Backend: backend, Force: force} - reply, err := messaging.RequestJSON[messaging.BackendStopRequest, messaging.BackendStopReply]( + reply, err := controlRequestJSON[workerctl.BackendStopRequest, workerctl.BackendStopReply]( a.nats, subject, req, backendStopAckTimeout) if err != nil { - if errors.Is(err, nats.ErrTimeout) { + if isStrictRequestTimeout(err) { xlog.Warn("Worker did not acknowledge backend.stop; assuming an older worker delivered it", "nodeID", nodeID, "backend", backend, "force", force) return nil @@ -490,11 +490,11 @@ func (a *RemoteUnloaderAdapter) stopBackend(nodeID, backend string, force bool) } // DeleteBackend tells a worker node to delete a backend (stop + remove files). -func (a *RemoteUnloaderAdapter) DeleteBackend(nodeID, backendName string) (*messaging.BackendDeleteReply, error) { +func (a *RemoteUnloaderAdapter) DeleteBackend(nodeID, backendName string) (*workerctl.BackendDeleteReply, error) { subject := messaging.SubjectNodeBackendDelete(nodeID) xlog.Info("Sending NATS backend.delete", "nodeID", nodeID, "backend", backendName) - reply, err := messaging.RequestJSON[messaging.BackendDeleteRequest, messaging.BackendDeleteReply](a.nats, subject, messaging.BackendDeleteRequest{Backend: backendName}, 2*time.Minute) + reply, err := controlRequestJSON[workerctl.BackendDeleteRequest, workerctl.BackendDeleteReply](a.nats, subject, workerctl.BackendDeleteRequest{Backend: backendName}, 2*time.Minute) if err != nil { return reply, err } @@ -551,7 +551,7 @@ func (a *RemoteUnloaderAdapter) UnloadModelOnNode(nodeID, modelName string) erro subject := messaging.SubjectNodeModelUnload(nodeID) xlog.Info("Sending NATS model.unload", "nodeID", nodeID, "model", modelName) - reply, err := messaging.RequestJSON[messaging.ModelUnloadRequest, messaging.ModelUnloadReply](a.nats, subject, messaging.ModelUnloadRequest{ModelName: modelName}, 30*time.Second) + reply, err := controlRequestJSON[workerctl.ModelUnloadRequest, workerctl.ModelUnloadReply](a.nats, subject, workerctl.ModelUnloadRequest{ModelName: modelName}, 30*time.Second) if err != nil { return err } @@ -574,7 +574,7 @@ func (a *RemoteUnloaderAdapter) DeleteModelFiles(modelName string) error { subject := messaging.SubjectNodeModelDelete(node.ID) xlog.Info("Sending NATS model.delete", "nodeID", node.ID, "model", modelName) - reply, err := messaging.RequestJSON[messaging.ModelDeleteRequest, messaging.ModelDeleteReply](a.nats, subject, messaging.ModelDeleteRequest{ModelName: modelName}, 30*time.Second) + reply, err := controlRequestJSON[workerctl.ModelDeleteRequest, workerctl.ModelDeleteReply](a.nats, subject, workerctl.ModelDeleteRequest{ModelName: modelName}, 30*time.Second) if err != nil { xlog.Warn("model.delete failed on node", "node", node.Name, "error", err) continue @@ -591,14 +591,3 @@ func (a *RemoteUnloaderAdapter) StopNode(nodeID string) error { subject := messaging.SubjectNodeStop(nodeID) return a.nats.Publish(subject, nil) } - -// isNATSTimeout returns true if err looks like a NATS request-reply timeout. -// nats.ErrTimeout is the canonical sentinel; context.DeadlineExceeded can -// also surface depending on the client's path; we accept both, plus a -// string-match fallback for clients that return a bare error. -func isNATSTimeout(err error) bool { - if errors.Is(err, nats.ErrTimeout) || errors.Is(err, context.DeadlineExceeded) { - return true - } - return err != nil && strings.Contains(err.Error(), "nats: timeout") -} diff --git a/core/services/nodes/unloader_ping_test.go b/core/services/nodes/unloader_ping_test.go index a9b3a5889..7231fbdb9 100644 --- a/core/services/nodes/unloader_ping_test.go +++ b/core/services/nodes/unloader_ping_test.go @@ -6,13 +6,13 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - "github.com/nats-io/nats.go" "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" ) // The scheduler's liveness probe asks a worker a question over NATS and treats -// "no responders" as proof the worker is gone. That is only sound if every +// "no route" as reason to skip the worker. That is only sound if every // worker in the fleet subscribes to the subject asked. // // It originally asked models.running, which arrived in 4.6. A 4.5 worker is @@ -35,10 +35,10 @@ var _ = Describe("Node liveness probe subject", func() { It("treats a worker that answers backend.list as alive", func() { // A worker old enough to predate models.running: it answers the // long-standing backend.list subject and nothing else. - mc.scriptReply(messaging.SubjectNodeBackendList(nodeID), messaging.BackendListReply{}) + mc.scriptReply(messaging.SubjectNodeBackendList(nodeID), workerctl.BackendListReply{}) mc.scriptNoResponders(messaging.SubjectNodeModelsRunning(nodeID)) - Expect(errors.Is(adapter.PingNode(nodeID), nats.ErrNoResponders)).To(BeFalse(), + Expect(errors.Is(adapter.PingNode(nodeID), ErrNoRoute)).To(BeFalse(), "a worker answering backend.list is alive regardless of newer subjects") }) @@ -46,6 +46,6 @@ var _ = Describe("Node liveness probe subject", func() { mc.scriptNoResponders(messaging.SubjectNodeBackendList(nodeID)) mc.scriptNoResponders(messaging.SubjectNodeModelsRunning(nodeID)) - Expect(errors.Is(adapter.PingNode(nodeID), nats.ErrNoResponders)).To(BeTrue()) + Expect(errors.Is(adapter.PingNode(nodeID), ErrNoRoute)).To(BeTrue()) }) }) diff --git a/core/services/nodes/unloader_stale_rows_test.go b/core/services/nodes/unloader_stale_rows_test.go index 0fff87baf..121e1aefe 100644 --- a/core/services/nodes/unloader_stale_rows_test.go +++ b/core/services/nodes/unloader_stale_rows_test.go @@ -4,10 +4,9 @@ import ( "encoding/json" "time" + "github.com/mudler/LocalAI/core/services/workerctl" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" - - "github.com/mudler/LocalAI/core/services/messaging" ) // Replies are handed to the adapter as raw JSON rather than as marshalled @@ -176,7 +175,7 @@ var _ = Describe("RemoteUnloaderAdapter stale replica rows", func() { // Guards the rolling-upgrade direction that matters: a new // controller must keep working against every worker already // deployed, not just ones rebuilt from this commit. - var reply messaging.BackendDeleteReply + var reply workerctl.BackendDeleteReply Expect(json.Unmarshal([]byte(`{"success": true}`), &reply)).To(Succeed()) Expect(reply.Success).To(BeTrue()) Expect(reply.ReportsStoppedProcesses).To(BeFalse()) diff --git a/core/services/nodes/unloader_test.go b/core/services/nodes/unloader_test.go index 3564f93c8..7efd39292 100644 --- a/core/services/nodes/unloader_test.go +++ b/core/services/nodes/unloader_test.go @@ -14,6 +14,7 @@ import ( "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" ) // --- Fakes --- @@ -145,7 +146,7 @@ var _ = Describe("RemoteUnloaderAdapter", func() { // backend.stop is request-reply, so the default fake must answer the // way a current worker does. Specs that care about the reply override // requestReply themselves. - mc.requestReply = mustJSON(messaging.BackendStopReply{ + mc.requestReply = mustJSON(workerctl.BackendStopReply{ Success: true, StoppedProcessKeys: []string{"llama#0"}, ReportsStoppedProcesses: true, @@ -253,9 +254,9 @@ var _ = Describe("RemoteUnloaderAdapter", func() { locator.nodes = []BackendNode{{ID: "node-1", Name: "worker-1"}} Expect(adapter.UnloadRemoteModelContext(context.Background(), "llama", true)).To(Succeed()) - var payload messaging.BackendStopRequest + var payload workerctl.BackendStopRequest Expect(json.Unmarshal(mc.requestCalls[0].Data, &payload)).To(Succeed()) - Expect(payload).To(Equal(messaging.BackendStopRequest{Backend: "llama", Force: true})) + Expect(payload).To(Equal(workerctl.BackendStopRequest{Backend: "llama", Force: true})) }) }) @@ -266,9 +267,9 @@ var _ = Describe("RemoteUnloaderAdapter", func() { Expect(mc.requestCalls[0].Subject).To(Equal(messaging.SubjectNodeBackendStop("node-1"))) // An empty Backend is the wire signal for "stop all"; the worker's - // decodeBackendStopRequest reads it the same way it read the bare + // decodeBackendStop reads it the same way it read the bare // nil payload this replaced. - var payload messaging.BackendStopRequest + var payload workerctl.BackendStopRequest Expect(json.Unmarshal(mc.requestCalls[0].Data, &payload)).To(Succeed()) Expect(payload.Backend).To(BeEmpty()) }) @@ -276,7 +277,7 @@ var _ = Describe("RemoteUnloaderAdapter", func() { // The bug this reply exists for: the worker could not stop what was // asked, and the caller was told everything was fine. It("reports a stop the worker could not carry out", func() { - mc.requestReply = mustJSON(messaging.BackendStopReply{ + mc.requestReply = mustJSON(workerctl.BackendStopReply{ Success: false, Error: "llama#0: process refused to die", ReportsStoppedProcesses: true, @@ -290,7 +291,7 @@ var _ = Describe("RemoteUnloaderAdapter", func() { // it stays a success — eviction and cleanup paths stop models that are // already gone all the time. It("succeeds when the worker matched no running process", func() { - mc.requestReply = mustJSON(messaging.BackendStopReply{ + mc.requestReply = mustJSON(workerctl.BackendStopReply{ Success: true, ReportsStoppedProcesses: true, }) @@ -316,7 +317,7 @@ var _ = Describe("RemoteUnloaderAdapter", func() { Expect(adapter.StopBackend("node-1", "llama-backend")).To(Succeed()) Expect(mc.requestCalls).To(HaveLen(1)) - var payload messaging.BackendStopRequest + var payload workerctl.BackendStopRequest Expect(json.Unmarshal(mc.requestCalls[0].Data, &payload)).To(Succeed()) Expect(payload.Backend).To(Equal("llama-backend")) Expect(payload.Force).To(BeFalse()) @@ -325,7 +326,7 @@ var _ = Describe("RemoteUnloaderAdapter", func() { Describe("StopModelReplica", func() { It("requests an acknowledged stop for the exact process", func() { - mc.requestReply, _ = json.Marshal(messaging.ModelStopReply{Matched: true, Terminated: true, ProcessKey: "llama#2"}) + mc.requestReply, _ = json.Marshal(workerctl.ModelStopReply{Matched: true, Terminated: true, ProcessKey: "llama#2"}) replica := NodeModel{ModelName: "llama", ReplicaIndex: 2, Address: "127.0.0.1:5002", ConfigRevision: "rev-1"} reply, err := adapter.StopModelReplica(context.Background(), "node-1", replica, true) @@ -335,9 +336,9 @@ var _ = Describe("RemoteUnloaderAdapter", func() { Expect(mc.requestCalls[0].Subject).To(Equal(messaging.SubjectNodeModelStop("node-1"))) Expect(mc.requestCalls[0].Timeout).To(BeNumerically(">", 0)) - var request messaging.ModelStopRequest + var request workerctl.ModelStopRequest Expect(json.Unmarshal(mc.requestCalls[0].Data, &request)).To(Succeed()) - Expect(request).To(Equal(messaging.ModelStopRequest{ + Expect(request).To(Equal(workerctl.ModelStopRequest{ ModelName: "llama", ProcessKey: "llama#2", ExpectedAddress: "127.0.0.1:5002", Force: true, ConfigRevision: "rev-1", })) }) @@ -427,7 +428,7 @@ func (f *failOnceMessagingClient) Close() {} var _ = Describe("RemoteUnloaderAdapter timeout configuration", func() { It("passes the configured install timeout to the messaging client", func() { mc := newScriptedMessagingClient() - mc.scriptReply(messaging.SubjectNodeBackendInstall("n1"), messaging.BackendInstallReply{Success: true, Address: "127.0.0.1:0"}) + mc.scriptReply(messaging.SubjectNodeBackendInstall("n1"), workerctl.BackendInstallReply{Success: true, Address: "127.0.0.1:0"}) adapter := NewRemoteUnloaderAdapter(nil, mc, 7*time.Minute, 11*time.Minute) _, err := adapter.InstallBackend("n1", "llama-cpp", "", "[]", "", "", "", 0, "", nil) @@ -439,7 +440,7 @@ var _ = Describe("RemoteUnloaderAdapter timeout configuration", func() { It("passes the configured upgrade timeout to the messaging client", func() { mc := newScriptedMessagingClient() - mc.scriptReply(messaging.SubjectNodeBackendUpgrade("n1"), messaging.BackendUpgradeReply{Success: true}) + mc.scriptReply(messaging.SubjectNodeBackendUpgrade("n1"), workerctl.BackendUpgradeReply{Success: true}) adapter := NewRemoteUnloaderAdapter(nil, mc, 7*time.Minute, 11*time.Minute) _, err := adapter.UpgradeBackend("n1", "llama-cpp", "[]", "", "", "", 0, "", nil) @@ -470,25 +471,25 @@ var _ = Describe("RemoteUnloaderAdapter NATS timeout handling", func() { _, err := adapter.InstallBackend("n1", "vllm", "", "[]", "", "", "", 0, "", nil) Expect(err).To(HaveOccurred()) Expect(errors.Is(err, galleryop.ErrWorkerStillInstalling)).To(BeFalse()) - Expect(errors.Is(err, nats.ErrNoResponders)).To(BeTrue()) + Expect(errors.Is(err, ErrNoRoute)).To(BeTrue()) }) }) var _ = Describe("RemoteUnloaderAdapter install progress streaming", func() { It("forwards BackendInstallProgressEvent values into the onProgress callback when the worker publishes them", func() { mc := newScriptedMessagingClient() - mc.scriptReply(messaging.SubjectNodeBackendInstall("n1"), messaging.BackendInstallReply{Success: true, Address: "127.0.0.1:0"}) - mc.scheduleProgressPublish("n1", "op-abc", []messaging.BackendInstallProgressEvent{ + mc.scriptReply(messaging.SubjectNodeBackendInstall("n1"), workerctl.BackendInstallReply{Success: true, Address: "127.0.0.1:0"}) + mc.scheduleProgressPublish("n1", "op-abc", []workerctl.BackendInstallProgressEvent{ {OpID: "op-abc", NodeID: "n1", Backend: "vllm", FileName: "vllm.tar.zst", Current: "100 MB", Total: "1 GB", Percentage: 10}, {OpID: "op-abc", NodeID: "n1", Backend: "vllm", FileName: "vllm.tar.zst", Current: "500 MB", Total: "1 GB", Percentage: 50}, }) adapter := NewRemoteUnloaderAdapter(nil, mc, 1*time.Second, 1*time.Second) var ( - received []messaging.BackendInstallProgressEvent + received []workerctl.BackendInstallProgressEvent mu sync.Mutex ) - onProgress := func(ev messaging.BackendInstallProgressEvent) { + onProgress := func(ev workerctl.BackendInstallProgressEvent) { mu.Lock() defer mu.Unlock() received = append(received, ev) @@ -506,7 +507,7 @@ var _ = Describe("RemoteUnloaderAdapter install progress streaming", func() { It("does NOT subscribe when onProgress is nil (reconciler retry path)", func() { mc := newScriptedMessagingClient() - mc.scriptReply(messaging.SubjectNodeBackendInstall("n1"), messaging.BackendInstallReply{Success: true}) + mc.scriptReply(messaging.SubjectNodeBackendInstall("n1"), workerctl.BackendInstallReply{Success: true}) adapter := NewRemoteUnloaderAdapter(nil, mc, 1*time.Second, 1*time.Second) _, err := adapter.InstallBackend("n1", "vllm", "", "[]", "", "", "", 0, "", nil) diff --git a/core/services/nodes/unloader_upgrade_test.go b/core/services/nodes/unloader_upgrade_test.go index bad8f9ed5..21dfc7300 100644 --- a/core/services/nodes/unloader_upgrade_test.go +++ b/core/services/nodes/unloader_upgrade_test.go @@ -8,6 +8,7 @@ import ( . "github.com/onsi/gomega" "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" ) var _ = Describe("RemoteUnloaderAdapter.UpgradeBackend", func() { @@ -16,7 +17,7 @@ var _ = Describe("RemoteUnloaderAdapter.UpgradeBackend", func() { nodeID := "node-x" mc.scriptReply(messaging.SubjectNodeBackendUpgrade(nodeID), - messaging.BackendUpgradeReply{Success: true}) + workerctl.BackendUpgradeReply{Success: true}) adapter := NewRemoteUnloaderAdapter(nil, mc, 3*time.Minute, 15*time.Minute) reply, err := adapter.UpgradeBackend(nodeID, "llama-cpp", `[{"name":"x"}]`, "", "", "", 0, "", nil) @@ -43,17 +44,17 @@ var _ = Describe("RemoteUnloaderAdapter.UpgradeBackend", func() { opID := "op-upgrade-1" mc.scriptReply(messaging.SubjectNodeBackendUpgrade(nodeID), - messaging.BackendUpgradeReply{Success: true}) + workerctl.BackendUpgradeReply{Success: true}) // The worker would publish these while force-reinstalling. The harness // replays them as soon as the adapter subscribes to the per-op subject. - mc.scheduleProgressPublish(nodeID, opID, []messaging.BackendInstallProgressEvent{ + mc.scheduleProgressPublish(nodeID, opID, []workerctl.BackendInstallProgressEvent{ {NodeID: nodeID, FileName: "llama-cpp.tar", Current: "10 MB", Total: "100 MB", Percentage: 10}, {NodeID: nodeID, FileName: "llama-cpp.tar", Current: "100 MB", Total: "100 MB", Percentage: 100}, }) var mu sync.Mutex - var got []messaging.BackendInstallProgressEvent - onProgress := func(ev messaging.BackendInstallProgressEvent) { + var got []workerctl.BackendInstallProgressEvent + onProgress := func(ev workerctl.BackendInstallProgressEvent) { mu.Lock() got = append(got, ev) mu.Unlock() diff --git a/core/services/quantization/service.go b/core/services/quantization/service.go index 011543205..b393feb89 100644 --- a/core/services/quantization/service.go +++ b/core/services/quantization/service.go @@ -74,7 +74,7 @@ func NewQuantizationService( appConfig *config.ApplicationConfig, modelLoader *model.ModelLoader, configLoader *config.ModelConfigLoader, - nats messaging.MessagingClient, + nats messaging.Broadcaster, store *distributed.QuantStore, ) *QuantizationService { s := &QuantizationService{ diff --git a/core/services/syncstate/syncstate.go b/core/services/syncstate/syncstate.go index 5aa69470f..fcf4b7964 100644 --- a/core/services/syncstate/syncstate.go +++ b/core/services/syncstate/syncstate.go @@ -38,7 +38,7 @@ type Store[K comparable, V any] interface { type Config[K comparable, V any] struct { Name string // subject namespace, e.g. "finetune.jobs" Key func(V) K // extract the key from a value - Nats messaging.MessagingClient // nil => standalone: in-memory only, no broadcast/subscribe + Nats messaging.Broadcaster // nil => standalone: in-memory only, no broadcast/subscribe Store Store[K, V] // optional read-through persistence Loader func(ctx context.Context) ([]V, error) // source when there is no Store (e.g. disk reload) OnApply func(op string, k K, v V) // optional hook after an applied change (e.g. ShutdownModel) @@ -111,7 +111,7 @@ func (m *SyncedMap[K, V]) Start(ctx context.Context) error { // nats.go transparently resubscribes on reconnect, but it cannot know we // kept derived in-memory state that may have drifted while the link was // down, so re-hydrate from the durable source. Detected via an optional - // interface so MessagingClient itself stays minimal; standalone/test + // interface so Broadcaster itself stays minimal; standalone/test // clients without the method simply fall back to the reconcile ticker. if r, ok := m.cfg.Nats.(interface{ OnReconnect(func()) }); ok { r.OnReconnect(func() { diff --git a/core/services/testutil/export_test.go b/core/services/testutil/export_test.go new file mode 100644 index 000000000..cc6da89af --- /dev/null +++ b/core/services/testutil/export_test.go @@ -0,0 +1,4 @@ +package testutil + +// SubjectMatches exposes the fake bus matching rule to the external specs. +var SubjectMatches = subjectMatches diff --git a/core/services/testutil/fakebus.go b/core/services/testutil/fakebus.go index 7452d810f..8d99d02fa 100644 --- a/core/services/testutil/fakebus.go +++ b/core/services/testutil/fakebus.go @@ -23,6 +23,9 @@ import ( type FakeBus struct { mu sync.Mutex subs []fakeBusSub + // nextID gives every subscription an identity of its own, so Unsubscribe + // removes that subscription and not another one on the same subject. + nextID uint64 // publishCounts records how many messages were published per subject, so a // spec can assert the echo-loop guard (an applied delta must not re-publish). publishCounts map[string]int @@ -31,43 +34,37 @@ type FakeBus struct { // spec exercise the component's reconnect re-hydrate path without a real // NATS server. reconnectCbs []func() + + // queueGroups records the queue group each queue subscription asked for, + // keyed by subject, because a group name decides which processes compete + // and a spec has to be able to pin it. + queueGroups map[string]string + // replyHandlers keeps each reply subscription's handler so a spec can play + // the requester through DeliverReply. + replyHandlers map[string]func([]byte, func([]byte)) } type fakeBusSub struct { + id uint64 subject string handler func([]byte) } // NewFakeBus returns a ready-to-use in-memory bus. func NewFakeBus() *FakeBus { - return &FakeBus{publishCounts: map[string]int{}} -} - -// subjectMatches reports whether a subscription filter matches a concrete -// subject, honoring the single-token `*` wildcard used by NATS. -func subjectMatches(filter, subject string) bool { - if filter == subject { - return true + return &FakeBus{ + publishCounts: map[string]int{}, + queueGroups: map[string]string{}, + replyHandlers: map[string]func([]byte, func([]byte)){}, } - fp := strings.Split(filter, ".") - sp := strings.Split(subject, ".") - if len(fp) != len(sp) { - return false - } - for i := range fp { - if fp[i] == "*" { - continue - } - if fp[i] != sp[i] { - return false - } - } - return true } // Publish marshals data as JSON and delivers it synchronously to every matching // subscriber. func (b *FakeBus) Publish(subject string, data any) error { + if err := messaging.ValidateSubject(subject); err != nil { + return err + } payload, err := json.Marshal(data) if err != nil { return err @@ -100,7 +97,7 @@ func (s *fakeBusSubscription) Unsubscribe() error { s.bus.mu.Lock() defer s.bus.mu.Unlock() for i, candidate := range s.bus.subs { - if candidate.subject == s.subRef.subject { + if candidate.id == s.subRef.id { s.bus.subs = append(s.bus.subs[:i], s.bus.subs[i+1:]...) return nil } @@ -109,21 +106,69 @@ func (s *fakeBusSubscription) Unsubscribe() error { } func (b *FakeBus) Subscribe(subject string, handler func([]byte)) (messaging.Subscription, error) { - sub := fakeBusSub{subject: subject, handler: handler} + if err := messaging.ValidateSubject(subject); err != nil { + return nil, err + } b.mu.Lock() + b.nextID++ + sub := fakeBusSub{id: b.nextID, subject: subject, handler: handler} b.subs = append(b.subs, sub) b.mu.Unlock() return &fakeBusSubscription{bus: b, subRef: sub}, nil } -func (b *FakeBus) QueueSubscribe(subject, _ string, handler func([]byte)) (messaging.Subscription, error) { - return b.Subscribe(subject, handler) +func (b *FakeBus) QueueSubscribe(subject, queue string, handler func([]byte)) (messaging.Subscription, error) { + sub, err := b.Subscribe(subject, handler) + if err != nil { + return nil, err + } + b.mu.Lock() + b.queueGroups[subject] = queue + b.mu.Unlock() + return sub, nil } -func (b *FakeBus) QueueSubscribeReply(string, string, func([]byte, func([]byte))) (messaging.Subscription, error) { +func (b *FakeBus) QueueSubscribeReply(subject, queue string, handler func([]byte, func([]byte))) (messaging.Subscription, error) { + if err := messaging.ValidateSubject(subject); err != nil { + return nil, err + } + b.mu.Lock() + b.queueGroups[subject] = queue + b.replyHandlers[subject] = handler + b.mu.Unlock() return &fakeBusSubscription{bus: b}, nil } +// QueueGroups returns a copy of the queue group recorded for each subject by +// QueueSubscribe and QueueSubscribeReply. +func (b *FakeBus) QueueGroups() map[string]string { + b.mu.Lock() + defer b.mu.Unlock() + out := make(map[string]string, len(b.queueGroups)) + for k, v := range b.queueGroups { + out[k] = v + } + return out +} + +// DeliverReply calls the reply handler registered on the exact subject with +// data and returns what it replied. ok is false when no handler is registered +// on the subject or the handler returned without replying, the two cases a +// real requester sees as a timeout. +func (b *FakeBus) DeliverReply(subject string, data []byte) (reply []byte, ok bool) { + b.mu.Lock() + h := b.replyHandlers[subject] + b.mu.Unlock() + if h == nil { + return nil, false + } + h(data, func(r []byte) { + reply = r + ok = true + }) + return reply, ok +} + func (b *FakeBus) SubscribeReply(string, func([]byte, func([]byte))) (messaging.Subscription, error) { return &fakeBusSubscription{bus: b}, nil } @@ -158,3 +203,26 @@ func (b *FakeBus) TriggerReconnect() { cb() } } + +// subjectMatches reports whether a subscription filter matches a concrete +// subject, honouring the single-token `*` wildcard the way NATS does, so the +// fake delivers to the same subscribers the real carrier would. +func subjectMatches(filter, subject string) bool { + if filter == subject { + return true + } + fp := strings.Split(filter, ".") + sp := strings.Split(subject, ".") + if len(fp) != len(sp) { + return false + } + for i := range fp { + if fp[i] == "*" { + continue + } + if fp[i] != sp[i] { + return false + } + } + return true +} diff --git a/core/services/testutil/fakebus_conformance_test.go b/core/services/testutil/fakebus_conformance_test.go new file mode 100644 index 000000000..05f40ebec --- /dev/null +++ b/core/services/testutil/fakebus_conformance_test.go @@ -0,0 +1,15 @@ +package testutil_test + +import ( + . "github.com/onsi/ginkgo/v2" + + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/messaging/messagingtest" + "github.com/mudler/LocalAI/core/services/testutil" +) + +var _ = Describe("FakeBus", func() { + messagingtest.RunBroadcasterConformance(func() (messaging.Broadcaster, func()) { + return testutil.NewFakeBus(), func() {} + }) +}) diff --git a/core/services/testutil/fakebus_queue_test.go b/core/services/testutil/fakebus_queue_test.go new file mode 100644 index 000000000..82bc8b43a --- /dev/null +++ b/core/services/testutil/fakebus_queue_test.go @@ -0,0 +1,39 @@ +package testutil_test + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/testutil" +) + +var _ = Describe("FakeBus queue helpers", func() { + It("records the queue group per subject", func() { + bus := testutil.NewFakeBus() + _, err := bus.QueueSubscribe("jobs.new", "workers", func([]byte) {}) + Expect(err).ToNot(HaveOccurred()) + _, err = bus.QueueSubscribe("agent.execute", "agent-workers", func([]byte) {}) + Expect(err).ToNot(HaveOccurred()) + + Expect(bus.QueueGroups()).To(Equal(map[string]string{ + "jobs.new": "workers", + "agent.execute": "agent-workers", + })) + }) + + It("keeps a reply handler so a spec can drive it", func() { + bus := testutil.NewFakeBus() + _, err := bus.QueueSubscribeReply("mcp.tools.execute", "agent-workers", func(data []byte, reply func([]byte)) { + reply(append([]byte("echo:"), data...)) + }) + Expect(err).ToNot(HaveOccurred()) + + out, ok := bus.DeliverReply("mcp.tools.execute", []byte("hi")) + Expect(ok).To(BeTrue()) + Expect(string(out)).To(Equal("echo:hi")) + Expect(bus.QueueGroups()).To(HaveKeyWithValue("mcp.tools.execute", "agent-workers")) + + _, ok = bus.DeliverReply("mcp.discovery", nil) + Expect(ok).To(BeFalse()) + }) +}) diff --git a/core/services/testutil/subject_match_test.go b/core/services/testutil/subject_match_test.go new file mode 100644 index 000000000..8ad5957b8 --- /dev/null +++ b/core/services/testutil/subject_match_test.go @@ -0,0 +1,21 @@ +package testutil_test + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/testutil" +) + +var _ = Describe("FakeBus subject matching", func() { + DescribeTable("matches like a NATS single-token wildcard", + func(filter, subject string, want bool) { + Expect(testutil.SubjectMatches(filter, subject)).To(Equal(want)) + }, + Entry("exact", "jobs.new", "jobs.new", true), + Entry("wildcard hit", "jobs.*.cancel", "jobs.abc.cancel", true), + Entry("wildcard wrong tail", "jobs.*.cancel", "jobs.abc.result", false), + Entry("wildcard does not span tokens", "jobs.*", "jobs.a.b", false), + Entry("length mismatch", "jobs.new", "jobs.new.extra", false), + ) +}) diff --git a/core/services/testutil/testutil_suite_test.go b/core/services/testutil/testutil_suite_test.go new file mode 100644 index 000000000..4b15ce5f2 --- /dev/null +++ b/core/services/testutil/testutil_suite_test.go @@ -0,0 +1,13 @@ +package testutil_test + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestTestutil(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Testutil test suite") +} diff --git a/core/services/worker/control_nats.go b/core/services/worker/control_nats.go new file mode 100644 index 000000000..cf38e6a02 --- /dev/null +++ b/core/services/worker/control_nats.go @@ -0,0 +1,99 @@ +package worker + +import ( + "context" + "fmt" + + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// natsControlServer serves control verbs on this node's NATS subjects. +type natsControlServer struct { + bus messaging.MessagingClient + nodeID string +} + +func newNATSControlServer(bus messaging.MessagingClient, nodeID string) *natsControlServer { + return &natsControlServer{bus: bus, nodeID: nodeID} +} + +func (n *natsControlServer) subject(v controlVerb) (string, error) { + switch v { + case verbBackendInstall: + return messaging.SubjectNodeBackendInstall(n.nodeID), nil + case verbBackendUpgrade: + return messaging.SubjectNodeBackendUpgrade(n.nodeID), nil + case verbBackendStop: + return messaging.SubjectNodeBackendStop(n.nodeID), nil + case verbBackendDelete: + return messaging.SubjectNodeBackendDelete(n.nodeID), nil + case verbBackendList: + return messaging.SubjectNodeBackendList(n.nodeID), nil + case verbModelsRunning: + return messaging.SubjectNodeModelsRunning(n.nodeID), nil + case verbModelUnload: + return messaging.SubjectNodeModelUnload(n.nodeID), nil + case verbModelStop: + return messaging.SubjectNodeModelStop(n.nodeID), nil + case verbModelDelete: + return messaging.SubjectNodeModelDelete(n.nodeID), nil + case verbNodeStop: + return messaging.SubjectNodeStop(n.nodeID), nil + case verbFilesEnsure: + return messaging.SubjectNodeFilesEnsure(n.nodeID), nil + case verbFilesStage: + return messaging.SubjectNodeFilesStage(n.nodeID), nil + case verbFilesTemp: + return messaging.SubjectNodeFilesTemp(n.nodeID), nil + case verbFilesListDir: + return messaging.SubjectNodeFilesListDir(n.nodeID), nil + case verbFilesRelease: + return messaging.SubjectNodeFilesRelease(n.nodeID), nil + } + return "", fmt.Errorf("no NATS subject for control verb %q", v) +} + +// handle runs h inside the subscription callback, so NATS delivers one request +// of the verb at a time, as it did before the verbs had a carrier seam. The +// undecodable error is dropped because reply already carries the typed +// refusal the requester expects. A panic is deliberately not recovered: the +// worker exits and goes unhealthy instead of leaving the requester to time out. +func (n *natsControlServer) handle(v controlVerb, h controlHandler) error { + subject, err := n.subject(v) + if err != nil { + return fmt.Errorf("serving %s: %w", v, err) + } + if _, err := n.bus.SubscribeReply(subject, func(data []byte, reply func([]byte)) { + if r, _ := h(context.Background(), data); r != nil { + replyJSON(reply, r) + } + }); err != nil { + return fmt.Errorf("serving %s: %w", v, err) + } + return nil +} + +// handleWithProgress spawns a goroutine per request so a multi-minute install +// does not hold up the next request on the same subscription. +func (n *natsControlServer) handleWithProgress(v controlVerb, h progressControlHandler) error { + subject, err := n.subject(v) + if err != nil { + return fmt.Errorf("serving %s: %w", v, err) + } + progress := func(ev workerctl.BackendInstallProgressEvent) { + // A lost progress event only delays the UI bar; the terminal reply is + // what the requester acts on. + _ = n.bus.Publish(messaging.SubjectNodeBackendInstallProgress(n.nodeID, ev.OpID), ev) + } + if _, err := n.bus.SubscribeReply(subject, func(data []byte, reply func([]byte)) { + go func() { + if r, _ := h(context.Background(), data, progress); r != nil { + replyJSON(reply, r) + } + }() + }); err != nil { + return fmt.Errorf("serving %s: %w", v, err) + } + return nil +} diff --git a/core/services/worker/control_nats_test.go b/core/services/worker/control_nats_test.go new file mode 100644 index 000000000..681dae63a --- /dev/null +++ b/core/services/worker/control_nats_test.go @@ -0,0 +1,374 @@ +package worker + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "log/slog" + "os" + "sync" + "syscall" + "time" + + "github.com/mudler/xlog" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" + "github.com/mudler/LocalAI/pkg/system" +) + +// recordingBus is a messaging.MessagingClient that keeps every subscription +// callback so a spec can deliver a request by hand and see what the worker +// answers. A Subscribe callback is stored behind a reply func it never calls, +// so a spec can assert "nothing was sent" the same way for both kinds. +type recordingBus struct { + mu sync.Mutex + subjects []string + handlers map[string]func([]byte, func([]byte)) + failOn map[string]error + publish []published +} + +type published struct { + subject string + payload any +} + +func newRecordingBus() *recordingBus { + return &recordingBus{handlers: map[string]func([]byte, func([]byte)){}, failOn: map[string]error{}} +} + +func (b *recordingBus) record(subject string, h func([]byte, func([]byte))) (messaging.Subscription, error) { + b.mu.Lock() + defer b.mu.Unlock() + if err := b.failOn[subject]; err != nil { + return nil, err + } + b.subjects = append(b.subjects, subject) + b.handlers[subject] = h + return releaseSubscription{}, nil +} + +func (b *recordingBus) Publish(subject string, payload any) error { + b.mu.Lock() + defer b.mu.Unlock() + b.publish = append(b.publish, published{subject: subject, payload: payload}) + return nil +} + +func (b *recordingBus) published() []published { + b.mu.Lock() + defer b.mu.Unlock() + return append([]published(nil), b.publish...) +} +func (b *recordingBus) Subscribe(subject string, h func([]byte)) (messaging.Subscription, error) { + return b.record(subject, func(data []byte, _ func([]byte)) { h(data) }) +} +func (b *recordingBus) QueueSubscribe(string, string, func([]byte)) (messaging.Subscription, error) { + return releaseSubscription{}, nil +} +func (b *recordingBus) QueueSubscribeReply(string, string, func([]byte, func([]byte))) (messaging.Subscription, error) { + return releaseSubscription{}, nil +} +func (b *recordingBus) SubscribeReply(subject string, h func([]byte, func([]byte))) (messaging.Subscription, error) { + return b.record(subject, h) +} +func (b *recordingBus) Request(string, []byte, time.Duration) ([]byte, error) { return nil, nil } +func (b *recordingBus) IsConnected() bool { return true } +func (b *recordingBus) Close() {} + +func (b *recordingBus) subscribed() []string { + b.mu.Lock() + defer b.mu.Unlock() + return append([]string(nil), b.subjects...) +} + +// deliver runs the subscription callback for subject the way the NATS client +// would and returns a channel that receives every reply it sends. +func (b *recordingBus) deliver(subject string, body []byte) <-chan string { + b.mu.Lock() + h := b.handlers[subject] + b.mu.Unlock() + Expect(h).NotTo(BeNil(), "no subscription for %s", subject) + replies := make(chan string, 4) + h(body, func(data []byte) { replies <- string(data) }) + return replies +} + +func newLifecycleTestSupervisor(sigCh chan<- os.Signal) *backendSupervisor { + ss, err := system.GetSystemState(system.WithBackendPath(GinkgoT().TempDir()), system.WithModelPath(GinkgoT().TempDir())) + Expect(err).NotTo(HaveOccurred()) + return &backendSupervisor{ + cfg: &Config{}, + nodeID: "n1", + systemState: ss, + sigCh: sigCh, + processes: map[string]*backendProcess{}, + } +} + +func registerLifecycleForTest(s *backendSupervisor, bus *recordingBus) error { + return s.registerLifecycleVerbs(newNATSControlServer(bus, s.nodeID)) +} + +const malformedBody = `{"backend":` + +var _ = Describe("Worker control verbs over NATS", func() { + var ( + bus *recordingBus + sigCh chan os.Signal + s *backendSupervisor + ) + + BeforeEach(func() { + bus = newRecordingBus() + sigCh = make(chan os.Signal, 1) + s = newLifecycleTestSupervisor(sigCh) + }) + + It("subscribes exactly the ten lifecycle subjects of the node", func() { + Expect(registerLifecycleForTest(s, bus)).To(Succeed()) + Expect(bus.subscribed()).To(ConsistOf( + messaging.SubjectNodeBackendInstall("n1"), + messaging.SubjectNodeBackendUpgrade("n1"), + messaging.SubjectNodeBackendStop("n1"), + messaging.SubjectNodeBackendDelete("n1"), + messaging.SubjectNodeBackendList("n1"), + messaging.SubjectNodeModelsRunning("n1"), + messaging.SubjectNodeModelUnload("n1"), + messaging.SubjectNodeModelStop("n1"), + messaging.SubjectNodeModelDelete("n1"), + messaging.SubjectNodeStop("n1"), + )) + }) + + DescribeTable("answers a malformed body with the verb's refusal bytes", + func(subject func(string) string, want string) { + Expect(registerLifecycleForTest(s, bus)).To(Succeed()) + replies := bus.deliver(subject("n1"), []byte(malformedBody)) + Eventually(replies).Should(Receive(Equal(want))) + Consistently(replies, 50*time.Millisecond).ShouldNot(Receive()) + }, + Entry("backend.install", messaging.SubjectNodeBackendInstall, + `{"success":false,"error":"invalid request: unexpected end of JSON input"}`), + Entry("backend.upgrade", messaging.SubjectNodeBackendUpgrade, + `{"success":false,"error":"invalid request: unexpected end of JSON input"}`), + Entry("backend.stop", messaging.SubjectNodeBackendStop, + `{"success":false,"error":"invalid request: decoding backend stop request: unexpected end of JSON input","reports_stopped_processes":true}`), + Entry("backend.delete", messaging.SubjectNodeBackendDelete, + `{"success":false,"error":"invalid request: unexpected end of JSON input"}`), + Entry("model.unload", messaging.SubjectNodeModelUnload, + `{"success":false,"error":"invalid request: unexpected end of JSON input"}`), + Entry("model.stop", messaging.SubjectNodeModelStop, + `{"matched":false,"freed":false,"terminated":false,"process_key":"","error":"invalid request: unexpected end of JSON input"}`), + Entry("model.delete", messaging.SubjectNodeModelDelete, + `{"success":false,"error":"invalid request"}`), + ) + + DescribeTable("still answers a malformed body on a verb that ignores its body", + func(subject func(string) string, want string) { + Expect(registerLifecycleForTest(s, bus)).To(Succeed()) + replies := bus.deliver(subject("n1"), []byte(malformedBody)) + Eventually(replies).Should(Receive(Equal(want))) + }, + Entry("backend.list", messaging.SubjectNodeBackendList, `{"backends":null}`), + Entry("models.running", messaging.SubjectNodeModelsRunning, `{"models":[]}`), + ) + + It("signals shutdown on node.stop without replying, and never blocks on a repeat", func() { + Expect(registerLifecycleForTest(s, bus)).To(Succeed()) + replies := bus.deliver(messaging.SubjectNodeStop("n1"), nil) + Expect(sigCh).To(Receive(Equal(syscall.SIGTERM))) + Expect(replies).NotTo(Receive()) + + sigCh <- syscall.SIGINT + replies = bus.deliver(messaging.SubjectNodeStop("n1"), nil) + Expect(replies).NotTo(Receive()) + Expect(sigCh).To(Receive(Equal(syscall.SIGINT))) + }) + + It("aborts registration on a subscribe error and names the verb", func() { + denied := errors.New("permissions violation") + bus.failOn[messaging.SubjectNodeBackendStop("n1")] = denied + + err := s.registerLifecycleVerbs(newNATSControlServer(bus, "n1")) + + Expect(err).To(MatchError(denied)) + Expect(err.Error()).To(HavePrefix("serving backend.stop: ")) + Expect(bus.subscribed()).NotTo(ContainElement(messaging.SubjectNodeBackendDelete("n1"))) + Expect(bus.subscribed()).NotTo(ContainElement(messaging.SubjectNodeStop("n1"))) + }) + + It("runs a unary verb inside the callback and a with-progress verb beside it", func() { + srv := newNATSControlServer(bus, "n1") + release := make(chan struct{}) + blocked := func() (any, error) { + <-release + return struct{}{}, nil + } + Expect(srv.handle(verbBackendList, func(_ context.Context, _ []byte) (any, error) { return blocked() })).To(Succeed()) + Expect(srv.handleWithProgress(verbBackendInstall, func(_ context.Context, _ []byte, _ progressSink) (any, error) { return blocked() })).To(Succeed()) + + installReturned := make(chan (<-chan string), 1) + go func() { installReturned <- bus.deliver(messaging.SubjectNodeBackendInstall("n1"), nil) }() + var installReplies <-chan string + Eventually(installReturned).Should(Receive(&installReplies)) + Expect(installReplies).NotTo(Receive()) + + listReturned := make(chan (<-chan string), 1) + go func() { listReturned <- bus.deliver(messaging.SubjectNodeBackendList("n1"), nil) }() + Consistently(listReturned, 100*time.Millisecond).ShouldNot(Receive()) + + close(release) + var listReplies <-chan string + Eventually(listReturned).Should(Receive(&listReplies)) + Expect(listReplies).To(Receive(Equal(`{}`))) + Eventually(installReplies).Should(Receive(Equal(`{}`))) + }) +}) + +// lockedBuffer lets a spec read what a handler goroutine logged. +type lockedBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *lockedBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *lockedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +var _ = Describe("Worker control verbs: install progress and malformed requests", func() { + var ( + bus *recordingBus + s *backendSupervisor + ) + + BeforeEach(func() { + bus = newRecordingBus() + s = newLifecycleTestSupervisor(make(chan os.Signal, 1)) + Expect(registerLifecycleForTest(s, bus)).To(Succeed()) + }) + + // emitTwo stands in for a gallery download that ticks twice. It reports + // whether the handler handed it a callback at all, which is how an install + // without an OpID stays silent. + emitTwo := func(onDownload func(file, current, total string, percentage float64)) bool { + if onDownload == nil { + return false + } + onDownload("backend.tar", "1 MB", "2 MB", 50) + onDownload("backend.tar", "2 MB", "2 MB", 100) + return true + } + + progressOn := func(subject string) []workerctl.BackendInstallProgressEvent { + var evs []workerctl.BackendInstallProgressEvent + for _, p := range bus.published() { + Expect(p.subject).To(Equal(subject)) + ev, ok := p.payload.(workerctl.BackendInstallProgressEvent) + Expect(ok).To(BeTrue(), "progress payload is %T", p.payload) + evs = append(evs, ev) + } + return evs + } + + expectTwoEvents := func(evs []workerctl.BackendInstallProgressEvent) { + Expect(evs).To(HaveLen(2)) + for _, ev := range evs { + Expect(ev.OpID).To(Equal("op1")) + Expect(ev.NodeID).To(Equal("n1")) + Expect(ev.Backend).To(Equal("vllm")) + Expect(ev.Phase).To(Equal(workerctl.PhaseDownloading)) + } + // The second tick lands inside the debounce window, so it only reaches + // the bus through the terminal flush that runs before the reply. + Expect(evs[0].Percentage).To(Equal(50.0)) + Expect(evs[1].Percentage).To(Equal(100.0)) + } + + It("publishes install progress on the per-op subject before replying", func() { + s.installFn = func(_ workerctl.BackendInstallRequest, _ bool, onDownload func(string, string, string, float64)) (string, error) { + emitTwo(onDownload) + return "127.0.0.1:50051", nil + } + body, err := json.Marshal(workerctl.BackendInstallRequest{Backend: "vllm", OpID: "op1"}) + Expect(err).NotTo(HaveOccurred()) + + var reply string + Eventually(bus.deliver(messaging.SubjectNodeBackendInstall("n1"), body)).Should(Receive(&reply)) + Expect(reply).To(ContainSubstring(`"success":true`)) + expectTwoEvents(progressOn(messaging.SubjectNodeBackendInstallProgress("n1", "op1"))) + }) + + It("publishes upgrade progress on the per-op subject before replying", func() { + s.upgradeFn = func(_ workerctl.BackendUpgradeRequest, onDownload func(string, string, string, float64)) ([]string, error) { + emitTwo(onDownload) + return nil, nil + } + body, err := json.Marshal(workerctl.BackendUpgradeRequest{Backend: "vllm", OpID: "op1"}) + Expect(err).NotTo(HaveOccurred()) + + var reply string + Eventually(bus.deliver(messaging.SubjectNodeBackendUpgrade("n1"), body)).Should(Receive(&reply)) + Expect(reply).To(ContainSubstring(`"success":true`)) + expectTwoEvents(progressOn(messaging.SubjectNodeBackendInstallProgress("n1", "op1"))) + }) + + It("reports no progress for an install without an OpID", func() { + gotCallback := make(chan bool, 1) + s.installFn = func(_ workerctl.BackendInstallRequest, _ bool, onDownload func(string, string, string, float64)) (string, error) { + gotCallback <- emitTwo(onDownload) + return "127.0.0.1:50051", nil + } + body, err := json.Marshal(workerctl.BackendInstallRequest{Backend: "vllm"}) + Expect(err).NotTo(HaveOccurred()) + + Eventually(bus.deliver(messaging.SubjectNodeBackendInstall("n1"), body)).Should(Receive()) + Expect(gotCallback).To(Receive(BeFalse())) + Expect(bus.published()).To(BeEmpty()) + }) + + Context("with a malformed request", func() { + var logs *lockedBuffer + + BeforeEach(func() { + logs = &lockedBuffer{} + handler := slog.NewTextHandler(logs, &slog.HandlerOptions{Level: slog.LevelWarn}) + xlog.SetLogger(xlog.NewLoggerWithHandler(handler, xlog.LogLevelWarn)) + }) + + AfterEach(func() { + // xlog has no getter for the package logger, so restore the + // default the suite starts with. + xlog.SetLogger(xlog.NewLogger(xlog.LogLevel("info"), "text")) + }) + + DescribeTable("leaves a warning that names the verb", + func(subject func(string) string, verb string) { + Eventually(bus.deliver(subject("n1"), []byte(malformedBody))).Should(Receive()) + Expect(logs.String()).To(And( + ContainSubstring(`msg="Ignoring malformed control request"`), + ContainSubstring("verb="+verb), + ContainSubstring("unexpected end of JSON input"), + )) + }, + Entry("backend.install", messaging.SubjectNodeBackendInstall, "backend.install"), + Entry("backend.upgrade", messaging.SubjectNodeBackendUpgrade, "backend.upgrade"), + Entry("backend.delete", messaging.SubjectNodeBackendDelete, "backend.delete"), + Entry("model.unload", messaging.SubjectNodeModelUnload, "model.unload"), + Entry("model.stop", messaging.SubjectNodeModelStop, "model.stop"), + Entry("model.delete", messaging.SubjectNodeModelDelete, "model.delete"), + ) + }) +}) diff --git a/core/services/worker/control_server.go b/core/services/worker/control_server.go new file mode 100644 index 000000000..c53037f6b --- /dev/null +++ b/core/services/worker/control_server.go @@ -0,0 +1,110 @@ +package worker + +import ( + "context" + "encoding/json" + + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// controlVerb names one control verb independent of the carrier that delivers +// it. The NATS server maps it onto a per-node subject; a tunnel server maps it +// onto a path. +type controlVerb string + +const ( + verbBackendInstall controlVerb = "backend.install" + verbBackendUpgrade controlVerb = "backend.upgrade" + verbBackendStop controlVerb = "backend.stop" + verbBackendDelete controlVerb = "backend.delete" + verbBackendList controlVerb = "backend.list" + verbModelsRunning controlVerb = "models.running" + verbModelUnload controlVerb = "model.unload" + verbModelStop controlVerb = "model.stop" + verbModelDelete controlVerb = "model.delete" + verbNodeStop controlVerb = "node.stop" + verbFilesEnsure controlVerb = "files.ensure" + verbFilesStage controlVerb = "files.stage" + verbFilesTemp controlVerb = "files.temp" + verbFilesListDir controlVerb = "files.listdir" + verbFilesRelease controlVerb = "files.release" +) + +// progressSink receives install progress while a long-running verb runs. +type progressSink func(workerctl.BackendInstallProgressEvent) + +// controlHandler answers one request. A nil reply means the verb sends no +// answer (node.stop). undecodable is non-nil only when body could not be read +// as the verb's request; reply then holds the verb's typed refusal. The NATS +// server sends reply either way, which is today's behaviour. A carrier that can +// signal a malformed request out of band (HTTP 400) may send that instead. +// Only tests read undecodable today; it is kept as the hook for such a carrier. +type controlHandler func(ctx context.Context, body []byte) (reply any, undecodable error) + +// progressControlHandler is controlHandler for a verb that may run for minutes +// and reports progress while it runs. progress is never nil. +type progressControlHandler func(ctx context.Context, body []byte, progress progressSink) (reply any, undecodable error) + +// controlServer is the carrier the worker serves its control verbs on. A +// registration returns once the carrier will deliver requests for the verb, or +// an error that names the verb; a carrier-side refusal (a NATS permission +// violation) is an error here, never a silent no-op. handle may deliver +// requests concurrently (the NATS server happens to serialise per verb). +// handleWithProgress delivers each request on its own goroutine, because a verb +// that runs for minutes must not hold up the next request of the same verb. +type controlServer interface { + handle(verb controlVerb, h controlHandler) error + handleWithProgress(verb controlVerb, h progressControlHandler) error +} + +// unary types a controlHandler. Go interfaces cannot carry generic methods, so +// the typing lives in these adapters and the interface stays byte-level. +func unary[Req, Reply any](decode func([]byte) (Req, error), refuse func(error) Reply, h func(context.Context, Req) Reply) controlHandler { + return func(ctx context.Context, body []byte) (any, error) { + req, err := decode(body) + if err != nil { + return refuse(err), err + } + return h(ctx, req), nil + } +} + +// withProgress is unary for progressControlHandler. +func withProgress[Req, Reply any](decode func([]byte) (Req, error), refuse func(error) Reply, h func(context.Context, Req, progressSink) Reply) progressControlHandler { + return func(ctx context.Context, body []byte, p progressSink) (any, error) { + req, err := decode(body) + if err != nil { + return refuse(err), err + } + return h(ctx, req, p), nil + } +} + +// noReply builds the handler of a verb that has no request and no reply. +func noReply(h func(context.Context)) controlHandler { + return func(ctx context.Context, _ []byte) (any, error) { + h(ctx) + return nil, nil + } +} + +func decodeJSON[Req any](body []byte) (Req, error) { + var req Req + err := json.Unmarshal(body, &req) + return req, err +} + +// ignoreBody is the decode of a verb that never read its body (backend.list, +// models.running, files.temp). It never fails, so a malformed body is still +// answered, as today. +func ignoreBody[Req any]([]byte) (Req, error) { + var req Req + return req, nil +} + +// refuseNever is the refusal of a verb whose decode cannot fail (ignoreBody), +// so it is never called. +func refuseNever[Reply any](error) Reply { + var reply Reply + return reply +} diff --git a/core/services/worker/file_staging.go b/core/services/worker/file_staging.go index 3eea3a9e6..3dc6a21dc 100644 --- a/core/services/worker/file_staging.go +++ b/core/services/worker/file_staging.go @@ -2,7 +2,6 @@ package worker import ( "context" - "encoding/json" "errors" "fmt" "os" @@ -11,8 +10,8 @@ import ( "strings" "time" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/storage" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/safefile" "github.com/mudler/xlog" "golang.org/x/sync/singleflight" @@ -41,8 +40,13 @@ func isPathAllowed(path string, allowedDirs []string) bool { return false } -// subscribeFileStaging subscribes to NATS file staging subjects for this node. -func (cfg *Config) subscribeFileStaging(natsClient messaging.MessagingClient, nodeID string, capacity *EphemeralCapacityGuard) error { +// invalidFileRequest is the refusal every body-reading file verb sends for a +// body it cannot decode; the frontend matches on the error text only. +const invalidFileRequest = "invalid request" + +// registerFileStagingVerbs serves the file staging verbs, backed by the +// configured object storage. +func (cfg *Config) registerFileStagingVerbs(srv controlServer, capacity *EphemeralCapacityGuard) error { // Create FileManager with same S3 config as the frontend // TODO: propagate a caller-provided context once Config carries one s3Store, err := storage.NewS3Store(context.Background(), storage.S3Config{ @@ -62,178 +66,162 @@ func (cfg *Config) subscribeFileStaging(natsClient messaging.MessagingClient, no if err != nil { return fmt.Errorf("initializing file manager: %w", err) } - if err := subscribeFileReleaseWithCapacity(natsClient, nodeID, fm, cacheDir, capacity); err != nil { + if err := registerFileReleaseVerb(srv, fm, cacheDir, capacity); err != nil { return err } - var ensureGroup singleflight.Group - // Subscribe: files.ensure — download S3 key to local, reply with local path - if _, err := natsClient.SubscribeReply(messaging.SubjectNodeFilesEnsure(nodeID), func(data []byte, reply func([]byte)) { - var req struct { - Key string `json:"key"` - } - if err := json.Unmarshal(data, &req); err != nil { - replyJSON(reply, map[string]string{"error": "invalid request"}) - return - } - - value, err, _ := ensureGroup.Do(req.Key, func() (any, error) { - return ensureWorkerFile(context.Background(), fm, capacity, req.Key) - }) - if err != nil { - xlog.Error("File ensure failed", "key", req.Key, "error", err) - replyJSON(reply, map[string]string{"error": err.Error()}) - return - } - localPath, ok := value.(string) - if !ok { - replyJSON(reply, map[string]string{"error": fmt.Sprintf("unexpected file ensure result %T", value)}) - return - } - - xlog.Debug("File ensured locally", "key", req.Key, "path", localPath) - replyJSON(reply, map[string]string{"local_path": localPath}) - }); err != nil { - return fmt.Errorf("subscribing to files.ensure events: %w", err) + v := &fileStagingVerbs{cfg: cfg, fm: fm, cacheDir: cacheDir, capacity: capacity} + if err := srv.handle(verbFilesEnsure, unary(decodeJSON[workerctl.FileEnsureRequest], func(error) workerctl.FileEnsureReply { + return workerctl.FileEnsureReply{Error: invalidFileRequest} + }, v.ensure)); err != nil { + return err + } + if err := srv.handle(verbFilesStage, unary(decodeJSON[workerctl.FileStageRequest], func(error) workerctl.FileStageReply { + return workerctl.FileStageReply{Error: invalidFileRequest} + }, v.stage)); err != nil { + return err + } + if err := srv.handle(verbFilesTemp, unary(ignoreBody[workerctl.FileTempRequest], refuseNever[workerctl.FileTempReply], v.temp)); err != nil { + return err + } + if err := srv.handle(verbFilesListDir, unary(decodeJSON[workerctl.FileListDirRequest], func(error) workerctl.FileListDirReply { + return workerctl.FileListDirReply{Error: invalidFileRequest} + }, v.listDir)); err != nil { + return err } - // Subscribe: files.stage — upload local path to S3, reply with key - if _, err := natsClient.SubscribeReply(messaging.SubjectNodeFilesStage(nodeID), func(data []byte, reply func([]byte)) { - var req struct { - LocalPath string `json:"local_path"` - Key string `json:"key"` - } - if err := json.Unmarshal(data, &req); err != nil { - replyJSON(reply, map[string]string{"error": "invalid request"}) - return - } - - allowedDirs := []string{cacheDir} - if cfg.ModelsPath != "" { - allowedDirs = append(allowedDirs, cfg.ModelsPath) - } - if !isPathAllowed(req.LocalPath, allowedDirs) { - replyJSON(reply, map[string]string{"error": "path outside allowed directories"}) - return - } - - if err := fm.Upload(context.Background(), req.Key, req.LocalPath); err != nil { - xlog.Error("File stage failed", "path", req.LocalPath, "key", req.Key, "error", err) - replyJSON(reply, map[string]string{"error": err.Error()}) - return - } - - xlog.Debug("File staged to S3", "path", req.LocalPath, "key", req.Key) - replyJSON(reply, map[string]string{"key": req.Key}) - }); err != nil { - return fmt.Errorf("subscribing to files.stage events: %w", err) - } - - // Subscribe: files.temp — allocate temp file, reply with local path - if _, err := natsClient.SubscribeReply(messaging.SubjectNodeFilesTemp(nodeID), func(data []byte, reply func([]byte)) { - tmpDir := filepath.Join(cacheDir, "staging-tmp") - if err := os.MkdirAll(tmpDir, 0750); err != nil { - replyJSON(reply, map[string]string{"error": fmt.Sprintf("creating temp dir: %v", err)}) - return - } - - f, err := os.CreateTemp(tmpDir, "localai-staging-*.tmp") - if err != nil { - replyJSON(reply, map[string]string{"error": fmt.Sprintf("creating temp file: %v", err)}) - return - } - localPath := f.Name() - if err := f.Close(); err != nil { - replyJSON(reply, map[string]string{"error": fmt.Sprintf("closing temp file: %v", err)}) - return - } - - xlog.Debug("Allocated temp file", "path", localPath) - replyJSON(reply, map[string]string{"local_path": localPath}) - }); err != nil { - return fmt.Errorf("subscribing to files.temp events: %w", err) - } - - // Subscribe: files.listdir — list files in a local directory, reply with relative paths - if _, err := natsClient.SubscribeReply(messaging.SubjectNodeFilesListDir(nodeID), func(data []byte, reply func([]byte)) { - var req struct { - KeyPrefix string `json:"key_prefix"` - } - if err := json.Unmarshal(data, &req); err != nil { - replyJSON(reply, map[string]any{"error": "invalid request"}) - return - } - - // Resolve key prefix to local directory - dirPath := filepath.Join(cacheDir, req.KeyPrefix) - if rel, ok := strings.CutPrefix(req.KeyPrefix, storage.ModelKeyPrefix); ok && cfg.ModelsPath != "" { - dirPath = filepath.Join(cfg.ModelsPath, rel) - } else if rel, ok := strings.CutPrefix(req.KeyPrefix, storage.DataKeyPrefix); ok { - dirPath = filepath.Join(cacheDir, "..", "data", rel) - } - - // Sanitize to prevent directory traversal via crafted key_prefix - dirPath = filepath.Clean(dirPath) - cleanCache := filepath.Clean(cacheDir) - cleanModels := filepath.Clean(cfg.ModelsPath) - cleanData := filepath.Clean(filepath.Join(cacheDir, "..", "data")) - if !(strings.HasPrefix(dirPath, cleanCache+string(filepath.Separator)) || - dirPath == cleanCache || - (cleanModels != "." && strings.HasPrefix(dirPath, cleanModels+string(filepath.Separator))) || - dirPath == cleanModels || - strings.HasPrefix(dirPath, cleanData+string(filepath.Separator)) || - dirPath == cleanData) { - replyJSON(reply, map[string]any{"error": "invalid key prefix"}) - return - } - - var files []string - if err := filepath.WalkDir(dirPath, func(path string, d os.DirEntry, err error) error { - if err != nil { - return err - } - if !d.IsDir() { - rel, err := filepath.Rel(dirPath, path) - if err != nil { - return err - } - files = append(files, rel) - } - return nil - }); err != nil { - xlog.Error("Failed to list staged files", "keyPrefix", req.KeyPrefix, "dirPath", dirPath, "error", err) - replyJSON(reply, map[string]any{"error": err.Error()}) - return - } - - xlog.Debug("Listed remote dir", "keyPrefix", req.KeyPrefix, "dirPath", dirPath, "fileCount", len(files)) - replyJSON(reply, map[string]any{"files": files}) - }); err != nil { - return fmt.Errorf("subscribing to files.listdir events: %w", err) - } - - xlog.Info("Subscribed to file staging NATS subjects", "nodeID", nodeID) + xlog.Info("Serving file staging verbs") return nil } -func subscribeFileRelease(natsClient messaging.MessagingClient, nodeID string, fm *storage.FileManager, cacheDir string) error { - return subscribeFileReleaseWithCapacity(natsClient, nodeID, fm, cacheDir, nil) +// fileStagingVerbs holds what the file staging verbs share for the lifetime +// of the worker. +type fileStagingVerbs struct { + cfg *Config + fm *storage.FileManager + cacheDir string + capacity *EphemeralCapacityGuard + // ensureGroup lives as long as the verbs so concurrent ensures of one key + // share a single download and a single capacity reservation. + ensureGroup singleflight.Group } -func subscribeFileReleaseWithCapacity(natsClient messaging.MessagingClient, nodeID string, fm *storage.FileManager, cacheDir string, capacity *EphemeralCapacityGuard) error { - if _, err := natsClient.SubscribeReply(messaging.SubjectNodeFilesRelease(nodeID), func(data []byte, reply func([]byte)) { - var req struct { - Key string `json:"key"` - RequestID string `json:"request_id"` +// ensure downloads an object storage key into the local cache. +func (v *fileStagingVerbs) ensure(ctx context.Context, req workerctl.FileEnsureRequest) workerctl.FileEnsureReply { + value, err, _ := v.ensureGroup.Do(req.Key, func() (any, error) { + return ensureWorkerFile(ctx, v.fm, v.capacity, req.Key) + }) + if err != nil { + xlog.Error("File ensure failed", "key", req.Key, "error", err) + return workerctl.FileEnsureReply{Error: err.Error()} + } + localPath, ok := value.(string) + if !ok { + return workerctl.FileEnsureReply{Error: fmt.Sprintf("unexpected file ensure result %T", value)} + } + + xlog.Debug("File ensured locally", "key", req.Key, "path", localPath) + return workerctl.FileEnsureReply{LocalPath: localPath} +} + +// stage uploads a local file to object storage. +func (v *fileStagingVerbs) stage(ctx context.Context, req workerctl.FileStageRequest) workerctl.FileStageReply { + allowedDirs := []string{v.cacheDir} + if v.cfg.ModelsPath != "" { + allowedDirs = append(allowedDirs, v.cfg.ModelsPath) + } + if !isPathAllowed(req.LocalPath, allowedDirs) { + return workerctl.FileStageReply{Error: "path outside allowed directories"} + } + + if err := v.fm.Upload(ctx, req.Key, req.LocalPath); err != nil { + xlog.Error("File stage failed", "path", req.LocalPath, "key", req.Key, "error", err) + return workerctl.FileStageReply{Error: err.Error()} + } + + xlog.Debug("File staged to S3", "path", req.LocalPath, "key", req.Key) + return workerctl.FileStageReply{Key: req.Key} +} + +// temp allocates an empty temporary file in the staging cache. +func (v *fileStagingVerbs) temp(context.Context, workerctl.FileTempRequest) workerctl.FileTempReply { + tmpDir := filepath.Join(v.cacheDir, "staging-tmp") + if err := os.MkdirAll(tmpDir, 0750); err != nil { + return workerctl.FileTempReply{Error: fmt.Sprintf("creating temp dir: %v", err)} + } + + f, err := os.CreateTemp(tmpDir, "localai-staging-*.tmp") + if err != nil { + return workerctl.FileTempReply{Error: fmt.Sprintf("creating temp file: %v", err)} + } + localPath := f.Name() + if err := f.Close(); err != nil { + return workerctl.FileTempReply{Error: fmt.Sprintf("closing temp file: %v", err)} + } + + xlog.Debug("Allocated temp file", "path", localPath) + return workerctl.FileTempReply{LocalPath: localPath} +} + +// listDir lists the files below a key prefix, relative to its directory. +func (v *fileStagingVerbs) listDir(_ context.Context, req workerctl.FileListDirRequest) workerctl.FileListDirReply { + cacheDir := v.cacheDir + // Resolve key prefix to local directory + dirPath := filepath.Join(cacheDir, req.KeyPrefix) + if rel, ok := strings.CutPrefix(req.KeyPrefix, storage.ModelKeyPrefix); ok && v.cfg.ModelsPath != "" { + dirPath = filepath.Join(v.cfg.ModelsPath, rel) + } else if rel, ok := strings.CutPrefix(req.KeyPrefix, storage.DataKeyPrefix); ok { + dirPath = filepath.Join(cacheDir, "..", "data", rel) + } + + // Sanitize to prevent directory traversal via crafted key_prefix + dirPath = filepath.Clean(dirPath) + cleanCache := filepath.Clean(cacheDir) + cleanModels := filepath.Clean(v.cfg.ModelsPath) + cleanData := filepath.Clean(filepath.Join(cacheDir, "..", "data")) + if !(strings.HasPrefix(dirPath, cleanCache+string(filepath.Separator)) || + dirPath == cleanCache || + (cleanModels != "." && strings.HasPrefix(dirPath, cleanModels+string(filepath.Separator))) || + dirPath == cleanModels || + strings.HasPrefix(dirPath, cleanData+string(filepath.Separator)) || + dirPath == cleanData) { + return workerctl.FileListDirReply{Error: "invalid key prefix"} + } + + var files []string + if err := filepath.WalkDir(dirPath, func(path string, d os.DirEntry, err error) error { + if err != nil { + return err } - if err := json.Unmarshal(data, &req); err != nil { - replyJSON(reply, map[string]string{"error": "invalid request"}) - return + if !d.IsDir() { + rel, err := filepath.Rel(dirPath, path) + if err != nil { + return err + } + files = append(files, rel) } + return nil + }); err != nil { + xlog.Error("Failed to list staged files", "keyPrefix", req.KeyPrefix, "dirPath", dirPath, "error", err) + return workerctl.FileListDirReply{Error: err.Error()} + } + + xlog.Debug("Listed remote dir", "keyPrefix", req.KeyPrefix, "dirPath", dirPath, "fileCount", len(files)) + return workerctl.FileListDirReply{Files: files} +} + +// registerFileReleaseVerb serves files.release, which evicts one exact +// ephemeral key or every key staged for one request. capacity may be nil. +func registerFileReleaseVerb(srv controlServer, fm *storage.FileManager, cacheDir string, capacity *EphemeralCapacityGuard) error { + return srv.handle(verbFilesRelease, unary(decodeJSON[workerctl.FileReleaseRequest], func(error) workerctl.FileReleaseReply { + return workerctl.FileReleaseReply{Error: invalidFileRequest} + }, func(ctx context.Context, req workerctl.FileReleaseRequest) workerctl.FileReleaseReply { var err error if req.RequestID != "" { - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - err = releaseEphemeralCacheRequest(ctx, cacheDir, req.RequestID, capacity) + // Beginning a request release can wait on the capacity guard; the + // bound keeps one stuck request from holding up the verb. + releaseCtx, cancel := context.WithTimeout(ctx, 30*time.Second) + err = releaseEphemeralCacheRequest(releaseCtx, cacheDir, req.RequestID, capacity) cancel() } else { cachePath, cacheErr := fm.CachePath(req.Key) @@ -243,14 +231,10 @@ func subscribeFileReleaseWithCapacity(natsClient messaging.MessagingClient, node } } if err != nil { - replyJSON(reply, map[string]string{"error": err.Error()}) - return + return workerctl.FileReleaseReply{Error: err.Error()} } - replyJSON(reply, map[string]string{}) - }); err != nil { - return fmt.Errorf("subscribing to files.release events: %w", err) - } - return nil + return workerctl.FileReleaseReply{} + })) } func releaseEphemeralCacheKey(cacheDir, key string) error { diff --git a/core/services/worker/file_staging_release_test.go b/core/services/worker/file_staging_release_test.go index 1019d475d..683eb4c75 100644 --- a/core/services/worker/file_staging_release_test.go +++ b/core/services/worker/file_staging_release_test.go @@ -113,7 +113,7 @@ var _ = Describe("Worker exact-key staging release", func() { localPath := filepath.Join(canonicalWorkerTempDir(), "input.wav") Expect(os.WriteFile(localPath, content, 0o600)).To(Succeed()) - stager := nodes.NewHTTPFileStager(func(string) (string, error) { return addr, nil }, "secret") + stager := nodes.NewHTTPFileStager(func(string) (string, error) { return addr, nil }, "secret", nodes.DirectWorkerNetDialer()) for range 2 { path, ensureErr := stager.EnsureRemote(context.Background(), "worker", localPath, key) Expect(ensureErr).NotTo(HaveOccurred()) @@ -344,7 +344,7 @@ var _ = Describe("Worker exact-key staging release", func() { Expect(err).NotTo(HaveOccurred()) client := &releaseMessagingClient{} - Expect(subscribeFileRelease(client, "node.one", fm, cacheDir)).To(Succeed()) + Expect(registerFileReleaseVerb(newNATSControlServer(client, "node.one"), fm, cacheDir, nil)).To(Succeed()) Expect(client.subject).To(Equal(messaging.SubjectNodeFilesRelease("node.one"))) request, err := json.Marshal(map[string]string{"key": "ephemeral/request-id/audio/input.wav"}) Expect(err).NotTo(HaveOccurred()) @@ -371,7 +371,7 @@ var _ = Describe("Worker exact-key staging release", func() { fm, err := storage.NewFileManager(nil, cacheDir) Expect(err).NotTo(HaveOccurred()) client := &releaseMessagingClient{} - Expect(subscribeFileRelease(client, "node.one", fm, cacheDir)).To(Succeed()) + Expect(registerFileReleaseVerb(newNATSControlServer(client, "node.one"), fm, cacheDir, nil)).To(Succeed()) request, err := json.Marshal(map[string]any{"request_id": "request-id"}) Expect(err).NotTo(HaveOccurred()) var response []byte @@ -391,7 +391,7 @@ var _ = Describe("Worker exact-key staging release", func() { fm, err := storage.NewFileManager(nil, cacheDir) Expect(err).NotTo(HaveOccurred()) client := &releaseMessagingClient{} - Expect(subscribeFileRelease(client, "node-1", fm, cacheDir)).To(Succeed()) + Expect(registerFileReleaseVerb(newNATSControlServer(client, "node-1"), fm, cacheDir, nil)).To(Succeed()) request, err := json.Marshal(map[string]string{"key": "models/model.gguf"}) Expect(err).NotTo(HaveOccurred()) diff --git a/core/services/worker/file_staging_verbs_test.go b/core/services/worker/file_staging_verbs_test.go new file mode 100644 index 000000000..765b052c0 --- /dev/null +++ b/core/services/worker/file_staging_verbs_test.go @@ -0,0 +1,157 @@ +package worker + +import ( + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/messaging" +) + +// newFakeS3 answers every upload with 200 and every other object request +// with 404, so the stage verb can succeed and the ensure verb can fail +// without any real object storage. +func newFakeS3() *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPut { + w.WriteHeader(http.StatusOK) + return + } + w.WriteHeader(http.StatusNotFound) + })) +} + +func registerFileVerbsForTest(cfg *Config, bus *recordingBus) error { + return cfg.registerFileStagingVerbs(newNATSControlServer(bus, "n1"), nil) +} + +var _ = Describe("Worker file staging verbs over NATS", func() { + var ( + bus *recordingBus + cfg *Config + cacheDir string + s3 *httptest.Server + ) + + BeforeEach(func() { + s3 = newFakeS3() + DeferCleanup(s3.Close) + root := canonicalWorkerTempDir() + cfg = &Config{ + ModelsPath: filepath.Join(root, "models"), + StorageURL: s3.URL, + StorageBucket: "bucket", + StorageAccessKey: "key", + StorageSecretKey: "secret", + } + Expect(os.MkdirAll(cfg.ModelsPath, 0750)).To(Succeed()) + cacheDir = filepath.Join(root, "cache") + bus = newRecordingBus() + Expect(registerFileVerbsForTest(cfg, bus)).To(Succeed()) + }) + + reply := func(subject func(string) string, body string) string { + GinkgoHelper() + var got string + Eventually(bus.deliver(subject("n1"), []byte(body))).Should(Receive(&got)) + return got + } + + It("subscribes the five file subjects of the node, release first", func() { + Expect(bus.subscribed()).To(Equal([]string{ + messaging.SubjectNodeFilesRelease("n1"), + messaging.SubjectNodeFilesEnsure("n1"), + messaging.SubjectNodeFilesStage("n1"), + messaging.SubjectNodeFilesTemp("n1"), + messaging.SubjectNodeFilesListDir("n1"), + })) + }) + + DescribeTable("answers a malformed body with the invalid request refusal", + func(subject func(string) string) { + Expect(reply(subject, malformedBody)).To(Equal(`{"error":"invalid request"}`)) + }, + Entry("files.release", messaging.SubjectNodeFilesRelease), + Entry("files.ensure", messaging.SubjectNodeFilesEnsure), + Entry("files.stage", messaging.SubjectNodeFilesStage), + Entry("files.listdir", messaging.SubjectNodeFilesListDir), + ) + + It("allocates a temp file even when the body is malformed", func() { + got := reply(messaging.SubjectNodeFilesTemp, malformedBody) + Expect(got).To(MatchRegexp(`^\{"local_path":"` + filepath.Join(cacheDir, "staging-tmp") + `/localai-staging-[0-9]+\.tmp"\}$`)) + }) + + It("answers a temp dir failure with its error", func() { + tmpDir := filepath.Join(cacheDir, "staging-tmp") + Expect(os.WriteFile(tmpDir, []byte("x"), 0640)).To(Succeed()) + Expect(reply(messaging.SubjectNodeFilesTemp, `{}`)).To(Equal( + fmt.Sprintf(`{"error":"creating temp dir: mkdir %s: not a directory"}`, tmpDir))) + }) + + It("ensures a cached key and answers its local path", func() { + path := filepath.Join(cacheDir, "models", "m.gguf") + Expect(os.MkdirAll(filepath.Dir(path), 0750)).To(Succeed()) + Expect(os.WriteFile(path, []byte("x"), 0640)).To(Succeed()) + Expect(reply(messaging.SubjectNodeFilesEnsure, `{"key":"models/m.gguf"}`)).To(Equal( + fmt.Sprintf(`{"local_path":%q}`, path))) + }) + + It("answers an ensure failure with only an error", func() { + Expect(reply(messaging.SubjectNodeFilesEnsure, `{"key":"models/missing.gguf"}`)).To( + MatchRegexp(`^\{"error":"downloading models/missing.gguf: .+"\}$`)) + }) + + It("stages an allowed path and answers its key", func() { + path := filepath.Join(cfg.ModelsPath, "m.gguf") + Expect(os.WriteFile(path, []byte("x"), 0640)).To(Succeed()) + Expect(reply(messaging.SubjectNodeFilesStage, + fmt.Sprintf(`{"local_path":%q,"key":"models/m.gguf"}`, path))).To(Equal(`{"key":"models/m.gguf"}`)) + }) + + It("refuses to stage a path outside the allowed directories", func() { + Expect(reply(messaging.SubjectNodeFilesStage, `{"local_path":"/etc/passwd","key":"k"}`)).To(Equal( + `{"error":"path outside allowed directories"}`)) + }) + + It("lists the files under a key prefix", func() { + dir := filepath.Join(cacheDir, "listing") + Expect(os.MkdirAll(filepath.Join(dir, "sub"), 0750)).To(Succeed()) + Expect(os.WriteFile(filepath.Join(dir, "a.txt"), []byte("x"), 0640)).To(Succeed()) + Expect(os.WriteFile(filepath.Join(dir, "sub", "b.txt"), []byte("x"), 0640)).To(Succeed()) + Expect(reply(messaging.SubjectNodeFilesListDir, `{"key_prefix":"listing"}`)).To(Equal( + `{"files":["a.txt","sub/b.txt"]}`)) + }) + + // Before the typed reply this was {"files":null}; the frontend decodes both + // to a nil Files slice. + It("answers an empty listing with an empty object", func() { + Expect(os.MkdirAll(filepath.Join(cacheDir, "empty"), 0750)).To(Succeed()) + Expect(reply(messaging.SubjectNodeFilesListDir, `{"key_prefix":"empty"}`)).To(Equal(`{}`)) + }) + + It("refuses a key prefix that escapes the staging directories", func() { + Expect(reply(messaging.SubjectNodeFilesListDir, `{"key_prefix":"../../../etc"}`)).To(Equal( + `{"error":"invalid key prefix"}`)) + }) + + It("answers a listing failure with its error", func() { + missing := filepath.Join(cacheDir, "missing") + Expect(reply(messaging.SubjectNodeFilesListDir, `{"key_prefix":"missing"}`)).To(Equal( + fmt.Sprintf(`{"error":"lstat %s: no such file or directory"}`, missing))) + }) + + It("answers a successful release with an empty object", func() { + Expect(reply(messaging.SubjectNodeFilesRelease, `{"request_id":"req-1"}`)).To(Equal(`{}`)) + }) + + It("answers a refused release with its error", func() { + Expect(reply(messaging.SubjectNodeFilesRelease, `{"key":"models/model.gguf"}`)).To(Equal( + `{"error":"release key \"models/model.gguf\" must identify one file below ephemeral/"}`)) + }) +}) diff --git a/core/services/worker/install.go b/core/services/worker/install.go index 122b5d266..46a4693fb 100644 --- a/core/services/worker/install.go +++ b/core/services/worker/install.go @@ -12,8 +12,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/gallery" "github.com/mudler/LocalAI/core/services/galleryop" - "github.com/mudler/LocalAI/core/services/messaging" - "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/xlog" ) @@ -54,13 +53,14 @@ func buildProcessKey(modelID, backend string, replicaIndex int) string { // 4. Find backend binary // 5. Start gRPC process on a new port // -// Returns the gRPC address of the backend process. +// Returns the gRPC address of the backend process. downloadCb receives the +// gallery download ticks; nil keeps the install silent. // // ProcessKey includes the replica index so a worker with MaxReplicasPerModel>1 // can host multiple processes for the same model on distinct ports. Old // controllers (no replica_index in the request) implicitly target replica 0, // which preserves single-replica behavior. -func (s *backendSupervisor) installBackend(req messaging.BackendInstallRequest, force bool) (string, error) { +func (s *backendSupervisor) installBackend(req workerctl.BackendInstallRequest, force bool, downloadCb func(file, current, total string, percentage float64)) (string, error) { processKey := buildProcessKey(req.ModelID, req.Backend, int(req.ReplicaIndex)) if !force { @@ -129,20 +129,6 @@ func (s *backendSupervisor) installBackend(req messaging.BackendInstallRequest, galleries = reqGalleries } - // When the master tagged this install with an OpID, stream the - // gallery download progress back to it on the per-op NATS subject. - // Old masters that omit OpID stay on the silent path so they keep - // working without changes. The publisher releases its mutex before - // every Publish so a slow link never stalls the download loop, and - // the deferred Flush guarantees a terminal-percentage event reaches - // the master even when the install errors out. - var downloadCb func(file, current, total string, percentage float64) - if req.OpID != "" && s.nats != nil { - publisher := nodes.NewDebouncedInstallProgressPublisher(s.nats, s.nodeID, req.OpID, req.Backend, installProgressDebounce) - downloadCb = publisher.OnDownload - defer publisher.Flush() - } - // On upgrade, run the gallery install path even if the binary already // exists on disk: findBackend would otherwise short-circuit and we'd // restart the same stale binary. The force flag passed to @@ -196,8 +182,8 @@ func (s *backendSupervisor) installBackend(req messaging.BackendInstallRequest, // It returns the process keys it terminated so the controller can drop the // NodeModel rows addressing them: an upgrade stops every process using the // binary and starts none back up, recycling their gRPC ports while the rows -// still point at those addresses. -func (s *backendSupervisor) upgradeBackend(req messaging.BackendUpgradeRequest) ([]string, error) { +// still point at those addresses. downloadCb is as for installBackend. +func (s *backendSupervisor) upgradeBackend(req workerctl.BackendUpgradeRequest, downloadCb func(file, current, total string, percentage float64)) ([]string, error) { // Stop every live process for this backend (peer replicas + the bare // processKey). Same logic as the force branch in installBackend. toStop := s.resolveProcessKeysForBackend(s.backendIdentity(req.Backend)) @@ -228,18 +214,6 @@ func (s *backendSupervisor) upgradeBackend(req messaging.BackendUpgradeRequest) galleries = reqGalleries } - // When the master tagged this upgrade with an OpID, stream gallery download - // progress back on the per-op subject (reused from install — an upgrade is a - // force-reinstall). Old masters omit OpID and stay on the silent path. The - // deferred Flush guarantees a terminal-percentage event even if the upgrade - // errors out, so the master's per-node bar never hangs mid-download. - var downloadCb func(file, current, total string, percentage float64) - if req.OpID != "" && s.nats != nil { - publisher := nodes.NewDebouncedInstallProgressPublisher(s.nats, s.nodeID, req.OpID, req.Backend, installProgressDebounce) - downloadCb = publisher.OnDownload - defer publisher.Flush() - } - if req.URI != "" { xlog.Info("Upgrading backend from external URI", "backend", req.Backend, "uri", req.URI) if err := galleryop.InstallExternalBackend( diff --git a/core/services/worker/lifecycle.go b/core/services/worker/lifecycle.go index f9be39c8c..86579c2fc 100644 --- a/core/services/worker/lifecycle.go +++ b/core/services/worker/lifecycle.go @@ -11,174 +11,212 @@ import ( "syscall" "github.com/mudler/LocalAI/core/gallery" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/workerctl" grpc "github.com/mudler/LocalAI/pkg/grpc" "github.com/mudler/xlog" ) -// subscribeLifecycleEvents wires every NATS subject this worker accepts to its -// per-event handler method. Each handler lives on *backendSupervisor below; -// keeping the dispatcher to a single line per subject makes adding a new -// subject a 2-line patch (one line here, one new method) instead of grafting -// onto a monolith. -func (s *backendSupervisor) subscribeLifecycleEvents() error { - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeBackendInstall(s.nodeID), s.handleBackendInstall); err != nil { - return fmt.Errorf("subscribing to backend install events: %w", err) +// registerLifecycleVerbs serves every lifecycle verb this worker accepts on +// srv. Each verb is one line here and one typed method below, so adding a verb +// does not graft onto a monolith. +func (s *backendSupervisor) registerLifecycleVerbs(srv controlServer) error { + reg := []func() error{ + func() error { + return srv.handleWithProgress(verbBackendInstall, withProgress(decodeJSON[workerctl.BackendInstallRequest], refuseInstall, s.serveInstall)) + }, + func() error { + return srv.handleWithProgress(verbBackendUpgrade, withProgress(decodeJSON[workerctl.BackendUpgradeRequest], refuseUpgrade, s.serveUpgrade)) + }, + func() error { + return srv.handle(verbBackendStop, unary(decodeBackendStop, refuseBackendStop, s.stopBackends)) + }, + func() error { + return srv.handle(verbBackendDelete, unary(decodeJSON[workerctl.BackendDeleteRequest], refuseDelete, s.deleteBackend)) + }, + func() error { + return srv.handle(verbBackendList, unary(ignoreBody[workerctl.BackendListRequest], refuseNever[workerctl.BackendListReply], s.backendList)) + }, + func() error { + return srv.handle(verbModelsRunning, unary(ignoreBody[workerctl.ModelsRunningRequest], refuseNever[workerctl.ModelsRunningReply], s.modelsRunning)) + }, + func() error { + return srv.handle(verbModelUnload, unary(decodeJSON[workerctl.ModelUnloadRequest], refuseUnload, s.unloadModel)) + }, + func() error { + return srv.handle(verbModelStop, unary(decodeJSON[workerctl.ModelStopRequest], refuseModelStop, s.stopModelExactCtx)) + }, + func() error { + return srv.handle(verbModelDelete, unary(decodeJSON[workerctl.ModelDeleteRequest], refuseModelDelete, s.deleteModel)) + }, + func() error { return srv.handle(verbNodeStop, noReply(s.signalNodeStop)) }, } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeBackendUpgrade(s.nodeID), s.handleBackendUpgrade); err != nil { - return fmt.Errorf("subscribing to backend upgrade events: %w", err) - } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeBackendStop(s.nodeID), s.handleBackendStop); err != nil { - return fmt.Errorf("subscribing to backend stop events: %w", err) - } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeBackendDelete(s.nodeID), s.handleBackendDelete); err != nil { - return fmt.Errorf("subscribing to backend delete events: %w", err) - } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeBackendList(s.nodeID), s.handleBackendList); err != nil { - return fmt.Errorf("subscribing to backend list events: %w", err) - } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeModelsRunning(s.nodeID), s.handleModelsRunning); err != nil { - return fmt.Errorf("subscribing to models running events: %w", err) - } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeModelUnload(s.nodeID), s.handleModelUnload); err != nil { - return fmt.Errorf("subscribing to model unload events: %w", err) - } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeModelStop(s.nodeID), s.handleModelStop); err != nil { - return fmt.Errorf("subscribing to model stop events: %w", err) - } - if _, err := s.nats.SubscribeReply(messaging.SubjectNodeModelDelete(s.nodeID), s.handleModelDelete); err != nil { - return fmt.Errorf("subscribing to model delete events: %w", err) - } - if _, err := s.nats.Subscribe(messaging.SubjectNodeStop(s.nodeID), s.handleNodeStop); err != nil { - return fmt.Errorf("subscribing to node stop events: %w", err) + for _, r := range reg { + if err := r(); err != nil { + return err + } } return nil } -func (s *backendSupervisor) handleModelStop(data []byte, reply func([]byte)) { - var req messaging.ModelStopRequest - if err := json.Unmarshal(data, &req); err != nil { - replyJSON(reply, messaging.ModelStopReply{Error: fmt.Sprintf("invalid request: %v", err)}) - return +// The refusals below are the replies each verb sent for an undecodable body +// before the verbs had a carrier seam. Requesters may match on them, so they +// are kept byte for byte, including model.delete omitting the cause. Each one +// logs, because the verbs log receipt only after a successful decode and a +// malformed request would otherwise leave no trace on the worker. + +func refuseInstall(err error) workerctl.BackendInstallReply { + xlog.Warn("Ignoring malformed control request", "verb", verbBackendInstall, "error", err) + return workerctl.BackendInstallReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} +} + +func refuseUpgrade(err error) workerctl.BackendUpgradeReply { + xlog.Warn("Ignoring malformed control request", "verb", verbBackendUpgrade, "error", err) + return workerctl.BackendUpgradeReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} +} + +func refuseBackendStop(err error) workerctl.BackendStopReply { + xlog.Error("Ignoring malformed NATS backend.stop event", "error", err) + return workerctl.BackendStopReply{ + Error: fmt.Sprintf("invalid request: %v", err), + ReportsStoppedProcesses: true, } - replyJSON(reply, s.stopModelExact(req)) } -// handleBackendInstall is the NATS callback for backend.install — install -// backend (idempotent: skips download if binary exists on disk) + start gRPC -// process (request-reply). +func refuseDelete(err error) workerctl.BackendDeleteReply { + xlog.Warn("Ignoring malformed control request", "verb", verbBackendDelete, "error", err) + return workerctl.BackendDeleteReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} +} + +func refuseUnload(err error) workerctl.ModelUnloadReply { + xlog.Warn("Ignoring malformed control request", "verb", verbModelUnload, "error", err) + return workerctl.ModelUnloadReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} +} + +func refuseModelStop(err error) workerctl.ModelStopReply { + xlog.Warn("Ignoring malformed control request", "verb", verbModelStop, "error", err) + return workerctl.ModelStopReply{Error: fmt.Sprintf("invalid request: %v", err)} +} + +func refuseModelDelete(err error) workerctl.ModelDeleteReply { + xlog.Warn("Ignoring malformed control request", "verb", verbModelDelete, "error", err) + return workerctl.ModelDeleteReply{Success: false, Error: "invalid request"} +} + +func (s *backendSupervisor) stopModelExactCtx(_ context.Context, req workerctl.ModelStopRequest) workerctl.ModelStopReply { + return s.stopModelExact(req) +} + +// serveInstall answers backend.install: install the backend (idempotent: skips +// download if binary exists on disk) and start its gRPC process. // -// Each request runs in its own goroutine so that a slow install on one -// backend does NOT head-of-line-block install requests for unrelated -// backends arriving on the same subscription. Per-backend serialization -// is provided by lockBackend so two requests targeting the same on-disk -// artifact don't race the gallery directory. -func (s *backendSupervisor) handleBackendInstall(data []byte, reply func([]byte)) { - go func() { - xlog.Info("Received NATS backend.install event") - var req messaging.BackendInstallRequest - if err := json.Unmarshal(data, &req); err != nil { - resp := messaging.BackendInstallReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} - replyJSON(reply, resp) - return - } +// The server runs each request on its own goroutine so that a slow install on +// one backend does NOT head-of-line-block install requests for unrelated +// backends. Per-backend serialization is provided by lockBackend so two +// requests targeting the same on-disk artifact don't race the gallery +// directory. +func (s *backendSupervisor) serveInstall(_ context.Context, req workerctl.BackendInstallRequest, progress progressSink) workerctl.BackendInstallReply { + xlog.Info("Received NATS backend.install event") + release := s.lockBackend(req.Backend) + defer release() + downloadCb, flush := s.downloadProgress(req.OpID, req.Backend, progress) + defer flush() - release := s.lockBackend(req.Backend) - defer release() + // req.Force=true is the legacy path used by pre-2026-05-08 masters + // that don't know about backend.upgrade. Honor it so a rolling + // update with new worker + old master keeps working; new masters + // send to backend.upgrade instead. + install := s.installFn + if install == nil { + install = s.installBackend + } + 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()} + } - // req.Force=true is the legacy path used by pre-2026-05-08 masters - // that don't know about backend.upgrade. Honor it so a rolling - // update with new worker + old master keeps working; new masters - // send to backend.upgrade instead. - addr, err := s.installBackend(req, req.Force) + advertiseAddr := addr + advAddr := s.cfg.advertiseAddr() + if advAddr != addr { + _, port, err := net.SplitHostPort(addr) if err != nil { - xlog.Error("Failed to install backend via NATS", "error", err) - resp := messaging.BackendInstallReply{Success: false, Error: err.Error()} - replyJSON(reply, resp) - return + xlog.Error("Failed to parse backend listen address; using it unchanged", "addr", addr, "error", err) + } else if advertiseHost, _, err := net.SplitHostPort(advAddr); err != nil { + xlog.Error("Failed to parse worker advertise address; using backend listen address", "addr", advAddr, "error", err) + } else { + advertiseAddr = net.JoinHostPort(advertiseHost, port) } - - advertiseAddr := addr - advAddr := s.cfg.advertiseAddr() - if advAddr != addr { - _, port, err := net.SplitHostPort(addr) - if err != nil { - xlog.Error("Failed to parse backend listen address; using it unchanged", "addr", addr, "error", err) - } else if advertiseHost, _, err := net.SplitHostPort(advAddr); err != nil { - xlog.Error("Failed to parse worker advertise address; using backend listen address", "addr", advAddr, "error", err) - } else { - advertiseAddr = net.JoinHostPort(advertiseHost, port) - } - } - resp := messaging.BackendInstallReply{Success: true, Address: advertiseAddr} - replyJSON(reply, resp) - }() + } + return workerctl.BackendInstallReply{Success: true, Address: advertiseAddr} } -// handleBackendUpgrade is the NATS callback for backend.upgrade — force-reinstall -// a backend (request-reply). Lives on its own subscription so a multi-minute -// download here does NOT block the install fast-path subscription on the same -// worker. -func (s *backendSupervisor) handleBackendUpgrade(data []byte, reply func([]byte)) { - go func() { - xlog.Info("Received NATS backend.upgrade event") - var req messaging.BackendUpgradeRequest - if err := json.Unmarshal(data, &req); err != nil { - resp := messaging.BackendUpgradeReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} - replyJSON(reply, resp) - return - } +// serveUpgrade answers backend.upgrade: force-reinstall a backend. It is its +// own verb so a multi-minute download here does NOT block the install +// fast-path on the same worker. +func (s *backendSupervisor) serveUpgrade(_ context.Context, req workerctl.BackendUpgradeRequest, progress progressSink) workerctl.BackendUpgradeReply { + xlog.Info("Received NATS backend.upgrade event") + release := s.lockBackend(req.Backend) + defer release() + downloadCb, flush := s.downloadProgress(req.OpID, req.Backend, progress) + defer flush() - release := s.lockBackend(req.Backend) - defer release() - - // stopped is meaningful even on the error paths: it lists processes - // already terminated (and ports already recycled) before the failure, so - // the controller must drop those rows regardless of the outcome. - stopped, err := s.upgradeBackend(req) - if err != nil { - xlog.Error("Failed to upgrade backend via NATS", "error", err) - replyJSON(reply, messaging.BackendUpgradeReply{ - Success: false, - Error: err.Error(), - StoppedProcessKeys: stopped, - ReportsStoppedProcesses: true, - }) - return - } - replyJSON(reply, messaging.BackendUpgradeReply{ - Success: true, + // stopped is meaningful even on the error paths: it lists processes + // already terminated (and ports already recycled) before the failure, so + // the controller must drop those rows regardless of the outcome. + upgrade := s.upgradeFn + if upgrade == nil { + upgrade = s.upgradeBackend + } + stopped, err := upgrade(req, downloadCb) + if err != nil { + xlog.Error("Failed to upgrade backend via NATS", "error", err) + return workerctl.BackendUpgradeReply{ + Success: false, + Error: err.Error(), StoppedProcessKeys: stopped, ReportsStoppedProcesses: true, - }) - }() + } + } + return workerctl.BackendUpgradeReply{ + Success: true, + StoppedProcessKeys: stopped, + ReportsStoppedProcesses: true, + } } -// handleBackendStop is the NATS callback for backend.stop — stop a specific -// backend process and report what it terminated. +// downloadProgress returns the gallery download callback for one install or +// upgrade and the flush the caller must defer. Requesters that send no OpID +// predate progress reporting and get a nil callback, so they see no events. +// The debounce and the terminal flush sit here, in the handler path, so every +// carrier behind progress forwards what it receives and sees the same bounded +// event rate. The flush runs before the reply, so the requester sees the +// terminal percentage even when the install fails. +func (s *backendSupervisor) downloadProgress(opID, backend string, progress progressSink) (func(file, current, total string, percentage float64), func()) { + if opID == "" { + return nil, func() {} + } + sink := nodes.NewDebouncedInstallProgressSink(progress, s.nodeID, opID, backend, installProgressDebounce) + return sink.OnDownload, sink.Flush +} + +// stopBackends answers backend.stop: stop a specific backend process (or all +// of them) and report what it terminated. // // The reply is what lets the controller tell a stop that worked from one that // matched nothing or failed. Callers that publish without a reply subject (an // older controller) still work: SubscribeReply drops the response. -func (s *backendSupervisor) handleBackendStop(data []byte, reply func([]byte)) { - req, stopAll, err := decodeBackendStopRequest(data) - if err != nil { - xlog.Error("Ignoring malformed NATS backend.stop event", "error", err) - replyJSON(reply, messaging.BackendStopReply{ - Error: fmt.Sprintf("invalid request: %v", err), - ReportsStoppedProcesses: true, - }) - return - } - if stopAll { +func (s *backendSupervisor) stopBackends(_ context.Context, req workerctl.BackendStopRequest) workerctl.BackendStopReply { + // Stop-all is exactly an empty Backend (an empty body decodes to that too), + // so it is derived here, not carried by the decoder. + if req.Backend == "" { xlog.Info("Received NATS backend.stop event (all)", "force", req.Force) stopped := s.stopAllBackends(req.Force) - replyJSON(reply, messaging.BackendStopReply{ + return workerctl.BackendStopReply{ Success: true, StoppedProcessKeys: stopped, ReportsStoppedProcesses: true, - }) - return + } } xlog.Info("Received NATS backend.stop event", "backend", req.Backend, "force", req.Force) // The identifier may be a backend name, a model name, or an exact @@ -198,7 +236,7 @@ func (s *backendSupervisor) handleBackendStop(data []byte, reply func([]byte)) { // failure: stopping a backend that is not running is the state the caller // asked for. The empty list is what tells the caller nothing matched, and // ReportsStoppedProcesses is what makes that emptiness trustworthy. - res := messaging.BackendStopReply{ + res := workerctl.BackendStopReply{ Success: len(failures) == 0, StoppedProcessKeys: stopped, ReportsStoppedProcesses: true, @@ -206,29 +244,26 @@ func (s *backendSupervisor) handleBackendStop(data []byte, reply func([]byte)) { if len(failures) > 0 { res.Error = strings.Join(failures, "; ") } - replyJSON(reply, res) + return res } -func decodeBackendStopRequest(data []byte) (messaging.BackendStopRequest, bool, error) { +// decodeBackendStop accepts an empty body because older controllers publish +// backend.stop with no payload to mean stop all; it decodes to an empty +// Backend, which is how stopBackends recognises stop-all. +func decodeBackendStop(data []byte) (workerctl.BackendStopRequest, error) { if len(data) == 0 { - return messaging.BackendStopRequest{}, true, nil + return workerctl.BackendStopRequest{}, nil } - var req messaging.BackendStopRequest + var req workerctl.BackendStopRequest if err := json.Unmarshal(data, &req); err != nil { - return messaging.BackendStopRequest{}, false, fmt.Errorf("decoding backend stop request: %w", err) + return workerctl.BackendStopRequest{}, fmt.Errorf("decoding backend stop request: %w", err) } - return req, req.Backend == "", nil + return req, nil } -// handleBackendDelete is the NATS callback for backend.delete — stop the -// backend process if running, then remove its files from disk (request-reply). -func (s *backendSupervisor) handleBackendDelete(data []byte, reply func([]byte)) { - var req messaging.BackendDeleteRequest - if err := json.Unmarshal(data, &req); err != nil { - resp := messaging.BackendDeleteReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} - replyJSON(reply, resp) - return - } +// deleteBackend answers backend.delete: stop the backend process if running, +// then remove its files from disk. +func (s *backendSupervisor) deleteBackend(_ context.Context, req workerctl.BackendDeleteRequest) workerctl.BackendDeleteReply { xlog.Info("Received NATS backend.delete event", "backend", req.Backend) // Resolve the backend's identity (concrete name + alias) BEFORE touching @@ -255,8 +290,8 @@ func (s *backendSupervisor) handleBackendDelete(data []byte, reply func([]byte)) // key is appended only after its process is confirmed gone, which is what // lets the controller trust the list on the partial-failure replies below. stopped := make([]string, 0, len(keys)) - deleteReply := func(success bool, errMsg string) messaging.BackendDeleteReply { - return messaging.BackendDeleteReply{ + deleteReply := func(success bool, errMsg string) workerctl.BackendDeleteReply { + return workerctl.BackendDeleteReply{ Success: success, Error: errMsg, StoppedProcessKeys: stopped, @@ -271,8 +306,7 @@ func (s *backendSupervisor) handleBackendDelete(data []byte, reply func([]byte)) // "backend deleted" while the process keeps serving requests. xlog.Error("Failed to stop backend process during delete; aborting delete", "backend", req.Backend, "processKey", key, "error", err) - replyJSON(reply, deleteReply(false, fmt.Sprintf("could not stop running process %s: %v", key, err))) - return + return deleteReply(false, fmt.Sprintf("could not stop running process %s: %v", key, err)) } stopped = append(stopped, key) } @@ -280,32 +314,28 @@ func (s *backendSupervisor) handleBackendDelete(data []byte, reply func([]byte)) // Delete the backend files if err := gallery.DeleteBackendFromSystem(s.systemState, req.Backend); err != nil { xlog.Warn("Failed to delete backend files", "backend", req.Backend, "error", err) - replyJSON(reply, deleteReply(false, err.Error())) - return + return deleteReply(false, err.Error()) } // Re-register backends after deletion if err := gallery.RegisterBackends(s.systemState, s.ml); err != nil { xlog.Error("Failed to refresh registered backends after deletion", "backend", req.Backend, "error", err) - replyJSON(reply, deleteReply(false, err.Error())) - return + return deleteReply(false, err.Error()) } - replyJSON(reply, deleteReply(true, "")) + return deleteReply(true, "") } -// handleBackendList is the NATS callback for backend.list — reply with the -// installed backends from this node's gallery (request-reply). -func (s *backendSupervisor) handleBackendList(data []byte, reply func([]byte)) { +// backendList answers backend.list with the installed backends from this +// node's gallery. +func (s *backendSupervisor) backendList(_ context.Context, _ workerctl.BackendListRequest) workerctl.BackendListReply { xlog.Info("Received NATS backend.list event") backends, err := gallery.ListSystemBackends(s.systemState) if err != nil { - resp := messaging.BackendListReply{Error: err.Error()} - replyJSON(reply, resp) - return + return workerctl.BackendListReply{Error: err.Error()} } - var infos []messaging.NodeBackendInfo + var infos []workerctl.NodeBackendInfo for name, b := range backends { // Drop synthetic alias rows: ListSystemBackends emits an entry // keyed by the alias name that re-uses the chosen concrete's @@ -319,7 +349,7 @@ func (s *backendSupervisor) handleBackendList(data []byte, reply func([]byte)) { if b.Metadata != nil && b.Metadata.Name != "" && name != b.Metadata.Name { continue } - info := messaging.NodeBackendInfo{ + info := workerctl.NodeBackendInfo{ Name: name, IsSystem: b.IsSystem, IsMeta: b.IsMeta, @@ -334,20 +364,13 @@ func (s *backendSupervisor) handleBackendList(data []byte, reply func([]byte)) { infos = append(infos, info) } - resp := messaging.BackendListReply{Backends: infos} - replyJSON(reply, resp) + return workerctl.BackendListReply{Backends: infos} } -// handleModelUnload is the NATS callback for model.unload — call gRPC Free() -// to release GPU memory without killing the backend process (request-reply). -func (s *backendSupervisor) handleModelUnload(data []byte, reply func([]byte)) { +// unloadModel answers model.unload: call gRPC Free() to release GPU memory +// without killing the backend process. +func (s *backendSupervisor) unloadModel(ctx context.Context, req workerctl.ModelUnloadRequest) workerctl.ModelUnloadReply { xlog.Info("Received NATS model.unload event") - var req messaging.ModelUnloadRequest - if err := json.Unmarshal(data, &req); err != nil { - resp := messaging.ModelUnloadReply{Success: false, Error: fmt.Sprintf("invalid request: %v", err)} - replyJSON(reply, resp) - return - } // Find the backend address for this model's backend type // The request includes an Address field if the router knows which process to target @@ -366,39 +389,29 @@ func (s *backendSupervisor) handleModelUnload(data []byte, reply func([]byte)) { // 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(context.Background(), workerBackendFreeTimeout) + 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() } - resp := messaging.ModelUnloadReply{Success: true} - replyJSON(reply, resp) + return workerctl.ModelUnloadReply{Success: true} } -// handleModelDelete is the NATS callback for model.delete — remove model -// files from disk (request-reply). -func (s *backendSupervisor) handleModelDelete(data []byte, reply func([]byte)) { +// deleteModel answers model.delete: remove model files from disk. +func (s *backendSupervisor) deleteModel(_ context.Context, req workerctl.ModelDeleteRequest) workerctl.ModelDeleteReply { xlog.Info("Received NATS model.delete event") - var req messaging.ModelDeleteRequest - if err := json.Unmarshal(data, &req); err != nil { - replyJSON(reply, messaging.ModelDeleteReply{Success: false, Error: "invalid request"}) - return - } - if err := gallery.DeleteStagedModelFiles(s.cfg.ModelsPath, req.ModelName); err != nil { xlog.Warn("Failed to delete model files", "model", req.ModelName, "error", err) - replyJSON(reply, messaging.ModelDeleteReply{Success: false, Error: err.Error()}) - return + return workerctl.ModelDeleteReply{Success: false, Error: err.Error()} } - - replyJSON(reply, messaging.ModelDeleteReply{Success: true}) + return workerctl.ModelDeleteReply{Success: true} } -// handleNodeStop is the NATS callback for node.stop — trigger the normal -// shutdown path via sigCh so deferred cleanup runs (fire-and-forget). -func (s *backendSupervisor) handleNodeStop(data []byte) { +// signalNodeStop answers node.stop: trigger the normal shutdown path via sigCh +// so deferred cleanup runs. It never replies. +func (s *backendSupervisor) signalNodeStop(_ context.Context) { xlog.Info("Received NATS stop event — signaling shutdown") select { case s.sigCh <- syscall.SIGTERM: diff --git a/core/services/worker/model_stop_test.go b/core/services/worker/model_stop_test.go index 345b61a30..65ddaf20d 100644 --- a/core/services/worker/model_stop_test.go +++ b/core/services/worker/model_stop_test.go @@ -7,12 +7,12 @@ import ( "net" "sync/atomic" - "github.com/mudler/LocalAI/core/services/messaging" process "github.com/mudler/go-processmanager" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" gogrpc "google.golang.org/grpc" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" ) @@ -42,14 +42,13 @@ func startModelStopProcess() *process.Process { return proc } -func requestModelStop(s *backendSupervisor, req messaging.ModelStopRequest) messaging.ModelStopReply { +func requestModelStop(s *backendSupervisor, req workerctl.ModelStopRequest) workerctl.ModelStopReply { data, err := json.Marshal(req) Expect(err).NotTo(HaveOccurred()) - var response []byte - s.handleModelStop(data, func(data []byte) { response = append([]byte(nil), data...) }) - var reply messaging.ModelStopReply - Expect(json.Unmarshal(response, &reply)).To(Succeed()) - return reply + reply, undecodable := unary(decodeJSON[workerctl.ModelStopRequest], refuseModelStop, s.stopModelExactCtx)(context.Background(), data) + Expect(undecodable).NotTo(HaveOccurred()) + Expect(reply).To(BeAssignableToTypeOf(workerctl.ModelStopReply{})) + return reply.(workerctl.ModelStopReply) } var _ = Describe("Acknowledged exact model stop", func() { @@ -64,9 +63,9 @@ var _ = Describe("Acknowledged exact model stop", func() { "model#1": other, }} - reply := requestModelStop(s, messaging.ModelStopRequest{ModelName: "model", ProcessKey: "model#0", ExpectedAddress: addr}) + reply := requestModelStop(s, workerctl.ModelStopRequest{ModelName: "model", ProcessKey: "model#0", ExpectedAddress: addr}) - Expect(reply).To(Equal(messaging.ModelStopReply{Matched: true, Freed: true, Terminated: true, ProcessKey: "model#0", Address: addr})) + Expect(reply).To(Equal(workerctl.ModelStopReply{Matched: true, Freed: true, Terminated: true, ProcessKey: "model#0", Address: addr})) Expect(backend.freeCalls.Load()).To(Equal(int32(1))) Expect(s.processes).To(HaveKeyWithValue("model#1", other)) Expect(s.processes).NotTo(HaveKey("model#0")) @@ -83,7 +82,7 @@ var _ = Describe("Acknowledged exact model stop", func() { }() s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{"model#0": {proc: proc, addr: "127.0.0.1:50051", port: 50051}}} - reply := requestModelStop(s, messaging.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: "127.0.0.1:50052"}) + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: "127.0.0.1:50052"}) Expect(reply.Matched).To(BeTrue()) Expect(reply.Terminated).To(BeFalse()) @@ -94,8 +93,8 @@ var _ = Describe("Acknowledged exact model stop", func() { It("treats an absent exact process key as idempotently terminated", func() { s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{}} - reply := requestModelStop(s, messaging.ModelStopRequest{ProcessKey: "missing#0", ExpectedAddress: "127.0.0.1:50051"}) - Expect(reply).To(Equal(messaging.ModelStopReply{Matched: false, Terminated: true, ProcessKey: "missing#0"})) + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "missing#0", ExpectedAddress: "127.0.0.1:50051"}) + Expect(reply).To(Equal(workerctl.ModelStopReply{Matched: false, Terminated: true, ProcessKey: "missing#0"})) }) It("reports Free failure but still terminates the process", func() { @@ -105,7 +104,7 @@ var _ = Describe("Acknowledged exact model stop", func() { proc := startModelStopProcess() s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{"model#0": {proc: proc, addr: addr, port: port}}} - reply := requestModelStop(s, messaging.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: addr}) + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: addr}) Expect(reply.Matched).To(BeTrue()) Expect(reply.Freed).To(BeFalse()) @@ -121,7 +120,7 @@ var _ = Describe("Acknowledged exact model stop", func() { proc := startModelStopProcess() s := &backendSupervisor{cfg: &Config{}, processes: map[string]*backendProcess{"model#0": {proc: proc, addr: addr, port: port}}} - reply := requestModelStop(s, messaging.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: addr, Force: true}) + reply := requestModelStop(s, workerctl.ModelStopRequest{ProcessKey: "model#0", ExpectedAddress: addr, Force: true}) Expect(reply.Matched).To(BeTrue()) Expect(reply.Freed).To(BeFalse()) diff --git a/core/services/worker/models_running.go b/core/services/worker/models_running.go index efc4700a0..1f4782f87 100644 --- a/core/services/worker/models_running.go +++ b/core/services/worker/models_running.go @@ -1,10 +1,11 @@ package worker import ( + "context" "strconv" "strings" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/xlog" ) @@ -33,11 +34,11 @@ func parseProcessKey(key string) (modelID string, replicaIndex int, ok bool) { // // Processes being stopped are excluded: they are alive but on their way out, // and reporting them would resurrect a replica the controller just released. -func (s *backendSupervisor) runningModels() []messaging.RunningModelInfo { +func (s *backendSupervisor) runningModels() []workerctl.RunningModelInfo { s.mu.Lock() defer s.mu.Unlock() - running := make([]messaging.RunningModelInfo, 0, len(s.processes)) + running := make([]workerctl.RunningModelInfo, 0, len(s.processes)) for key, bp := range s.processes { if bp == nil || bp.stopping || bp.proc == nil || !bp.proc.IsAlive() { continue @@ -47,7 +48,7 @@ func (s *backendSupervisor) runningModels() []messaging.RunningModelInfo { xlog.Warn("Skipping unparseable process key when reporting running models", "key", key) continue } - running = append(running, messaging.RunningModelInfo{ + running = append(running, workerctl.RunningModelInfo{ ModelID: modelID, ReplicaIndex: replicaIndex, Address: bp.addr, @@ -56,10 +57,10 @@ func (s *backendSupervisor) runningModels() []messaging.RunningModelInfo { return running } -// handleModelsRunning answers a models.running request with this worker's live +// modelsRunning answers a models.running request with this worker's live // process set. -func (s *backendSupervisor) handleModelsRunning(_ []byte, reply func([]byte)) { +func (s *backendSupervisor) modelsRunning(_ context.Context, _ workerctl.ModelsRunningRequest) workerctl.ModelsRunningReply { running := s.runningModels() xlog.Debug("Answering models.running", "nodeID", s.nodeID, "count", len(running)) - replyJSON(reply, messaging.ModelsRunningReply{Models: running}) + return workerctl.ModelsRunningReply{Models: running} } diff --git a/core/services/worker/replica_test.go b/core/services/worker/replica_test.go index d7340ff32..b3352eddf 100644 --- a/core/services/worker/replica_test.go +++ b/core/services/worker/replica_test.go @@ -3,7 +3,7 @@ package worker import ( "encoding/json" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" process "github.com/mudler/go-processmanager" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -185,27 +185,35 @@ var _ = Describe("Worker per-replica process keying", func() { }) Describe("backend.stop request decoding", func() { + // stopBackends treats an empty Backend as stop-all, so these pin what + // reaches it for each body shape. It("preserves the legacy empty-payload stop-all command", func() { - req, stopAll, err := decodeBackendStopRequest(nil) + req, err := decodeBackendStop(nil) Expect(err).NotTo(HaveOccurred()) - Expect(stopAll).To(BeTrue()) - Expect(req).To(Equal(messaging.BackendStopRequest{})) + Expect(req).To(Equal(workerctl.BackendStopRequest{})) }) It("preserves force for a structured stop-all command", func() { - data, err := json.Marshal(messaging.BackendStopRequest{Force: true}) + data, err := json.Marshal(workerctl.BackendStopRequest{Force: true}) Expect(err).NotTo(HaveOccurred()) - req, stopAll, err := decodeBackendStopRequest(data) + req, err := decodeBackendStop(data) Expect(err).NotTo(HaveOccurred()) - Expect(stopAll).To(BeTrue()) - Expect(req.Force).To(BeTrue()) + Expect(req).To(Equal(workerctl.BackendStopRequest{Force: true})) + }) + + It("decodes a named backend to that name", func() { + data, err := json.Marshal(workerctl.BackendStopRequest{Backend: "llama-cpp"}) + Expect(err).NotTo(HaveOccurred()) + + req, err := decodeBackendStop(data) + Expect(err).NotTo(HaveOccurred()) + Expect(req.Backend).To(Equal("llama-cpp")) }) It("rejects malformed JSON instead of treating it as stop-all", func() { - _, stopAll, err := decodeBackendStopRequest([]byte(`{"backend":`)) + _, err := decodeBackendStop([]byte(`{"backend":`)) Expect(err).To(MatchError(ContainSubstring("decoding backend stop request"))) - Expect(stopAll).To(BeFalse()) }) }) }) diff --git a/core/services/worker/supervisor.go b/core/services/worker/supervisor.go index 60754efcf..9a5a7c4d3 100644 --- a/core/services/worker/supervisor.go +++ b/core/services/worker/supervisor.go @@ -14,7 +14,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/gallery" - "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/workerctl" grpc "github.com/mudler/LocalAI/pkg/grpc" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/system" @@ -111,9 +111,14 @@ type backendSupervisor struct { systemState *system.SystemState galleries []config.Gallery nodeID string - nats messaging.MessagingClient sigCh chan<- os.Signal // send shutdown signal instead of os.Exit + // installFn and upgradeFn are the installers serveInstall and serveUpgrade + // run. nil means installBackend and upgradeBackend; specs set them to drive + // the verbs without a gallery. + installFn func(req workerctl.BackendInstallRequest, force bool, downloadCb func(file, current, total string, percentage float64)) (string, error) + upgradeFn func(req workerctl.BackendUpgradeRequest, downloadCb func(file, current, total string, percentage float64)) ([]string, error) + mu sync.Mutex processes map[string]*backendProcess // key: backend name nextPort int // next unhanded-out port; grows within [minPort, maxPort] @@ -867,8 +872,8 @@ func (s *backendSupervisor) stopBackendExact(key string, force bool) error { // stopModelExact implements the acknowledged controller-to-worker stop path. // The address check and stopping reservation are one critical section so a // stale controller request can never stop a replacement under the same key. -func (s *backendSupervisor) stopModelExact(req messaging.ModelStopRequest) messaging.ModelStopReply { - reply := messaging.ModelStopReply{ProcessKey: req.ProcessKey} +func (s *backendSupervisor) stopModelExact(req workerctl.ModelStopRequest) workerctl.ModelStopReply { + reply := workerctl.ModelStopReply{ProcessKey: req.ProcessKey} s.mu.Lock() bp, ok := s.processes[req.ProcessKey] diff --git a/core/services/worker/worker.go b/core/services/worker/worker.go index b40980462..60d469574 100644 --- a/core/services/worker/worker.go +++ b/core/services/worker/worker.go @@ -231,7 +231,6 @@ func Run(ctx *cliContext.Context, cfg *Config) error { systemState: systemState, galleries: galleries, nodeID: nodeID, - nats: natsClient, sigCh: sigCh, processes: make(map[string]*backendProcess), portAffinity: make(map[string]portOwnership), @@ -258,14 +257,15 @@ func Run(ctx *cliContext.Context, cfg *Config) error { }), )) - if err := supervisor.subscribeLifecycleEvents(); err != nil { + control := newNATSControlServer(natsClient, nodeID) + if err := supervisor.registerLifecycleVerbs(control); err != nil { nodes.ShutdownFileTransferServer(httpServer) return fmt.Errorf("subscribing to worker lifecycle events: %w", err) } - // Subscribe to file staging NATS subjects if S3 is configured + // Serve the file staging verbs only when S3 is configured if cfg.StorageURL != "" { - if err := cfg.subscribeFileStaging(natsClient, nodeID, ephemeralCapacity); err != nil { + if err := cfg.registerFileStagingVerbs(control, ephemeralCapacity); err != nil { nodes.ShutdownFileTransferServer(httpServer) return fmt.Errorf("subscribing to file staging subjects: %w", err) } diff --git a/core/services/workerctl/backend.go b/core/services/workerctl/backend.go new file mode 100644 index 000000000..a11c3fb62 --- /dev/null +++ b/core/services/workerctl/backend.go @@ -0,0 +1,164 @@ +package workerctl + +// BackendInstallRequest is the payload for a backend.install control request. +type BackendInstallRequest struct { + Backend string `json:"backend"` + ModelID string `json:"model_id,omitempty"` + BackendGalleries string `json:"backend_galleries,omitempty"` + // URI is set for external installs (OCI image, URL, or path). When non-empty + // the worker routes to InstallExternalBackend instead of the gallery lookup. + URI string `json:"uri,omitempty"` + Name string `json:"name,omitempty"` + Alias string `json:"alias,omitempty"` + // ReplicaIndex selects which slot on the worker this load occupies, so two + // concurrent backend.install requests for the same model land on distinct + // gRPC processes and ports. Workers older than this field treat it as 0 + // (single-replica behavior: no collision because the controller never + // asks for replica > 0 on a node whose MaxReplicasPerModel is 1). + ReplicaIndex int32 `json:"replica_index,omitempty"` + // Force is retained on the wire only for backward compatibility with + // pre-2026-05-08 masters that did not know about backend.upgrade. New + // callers MUST send to messaging.SubjectNodeBackendUpgrade instead. Workers continue + // to honor Force=true here so a rolling update with new master + old + // worker still works (the master's install fallback path also uses this + // when backend.upgrade finds no route to the worker). + Force bool `json:"force,omitempty"` + // OpID identifies the admin-side operation. When non-empty the worker + // publishes BackendInstallProgressEvent values to + // messaging.SubjectNodeBackendInstallProgress(nodeID, OpID) while the install is + // 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"` +} + +// BackendInstallReply is the response from a backend.install control request. +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"` +} + +// BackendUpgradeRequest is the payload for a backend.upgrade control request. +// It is intentionally a strict subset of BackendInstallRequest: there is no +// Force field because the upgrade subject IS the force semantics; no ModelID +// because upgrade is backend-scoped (it stops every replica using the binary +// before re-installing). Per-replica restart happens on the next routine load. +type BackendUpgradeRequest struct { + Backend string `json:"backend"` + BackendGalleries string `json:"backend_galleries,omitempty"` + URI string `json:"uri,omitempty"` + Name string `json:"name,omitempty"` + Alias string `json:"alias,omitempty"` + // ReplicaIndex is informational: upgrade stops all replicas regardless, + // but the field lets future per-replica metadata (e.g. progress reporting + // scoped to a slot) ride the same wire without a v3 type. + ReplicaIndex int32 `json:"replica_index,omitempty"` + // OpID identifies the admin-side operation. When non-empty the worker + // publishes BackendInstallProgressEvent values to + // messaging.SubjectNodeBackendInstallProgress(nodeID, OpID) while the force-reinstall + // runs, so the master can stream per-node progress for upgrades exactly as + // it already does for installs (an upgrade IS a force-reinstall, so the + // install-progress subject is reused rather than minting a new one; that adds no new + // NATS permission or rolling-update compat surface). Empty on legacy callers. + OpID string `json:"op_id,omitempty"` +} + +// BackendUpgradeReply mirrors BackendInstallReply minus Address: upgrade does +// not start a process, so there is no port to advertise. The subsequent +// routine load will re-bind via backend.install and learn the new address. +type BackendUpgradeReply struct { + Success bool `json:"success"` + Error string `json:"error,omitempty"` + + // StoppedProcessKeys / ReportsStoppedProcesses carry the same + // stale-row-invalidation contract as on BackendDeleteReply; an upgrade + // force-stops every process using the binary and starts none back up, so it + // recycles ports exactly the way a delete does. See that type for why the + // boolean is not redundant with an empty list. + StoppedProcessKeys []string `json:"stopped_process_keys,omitempty"` + ReportsStoppedProcesses bool `json:"reports_stopped_processes,omitempty"` +} + +// BackendListRequest is the payload for a backend.list control request. +type BackendListRequest struct{} + +// BackendListReply is the response from a backend.list control request. +type BackendListReply struct { + Backends []NodeBackendInfo `json:"backends"` + Error string `json:"error,omitempty"` +} + +// NodeBackendInfo describes a backend installed on a worker node. +type NodeBackendInfo struct { + Name string `json:"name"` + IsSystem bool `json:"is_system"` + IsMeta bool `json:"is_meta"` + InstalledAt string `json:"installed_at,omitempty"` + GalleryURL string `json:"gallery_url,omitempty"` + // Version, URI and Digest enable cluster-wide upgrade detection: + // without them, the frontend cannot tell whether the installed OCI + // image matches the gallery entry, and upgrades silently never surface. + Version string `json:"version,omitempty"` + URI string `json:"uri,omitempty"` + Digest string `json:"digest,omitempty"` +} + +// BackendStopRequest controls worker-side process shutdown. Force skips the +// best-effort Free RPC so a backend stuck serving a request can still be +// terminated by the watchdog. +type BackendStopRequest struct { + Backend string `json:"backend"` + Force bool `json:"force,omitempty"` +} + +// BackendStopReply is the worker's answer to a backend.stop request. +// +// backend.stop had no reply until this type existed. The controller published +// and returned success as soon as the local publish succeeded, so a stop that +// killed nothing, and a stop that failed outright, both looked identical to a +// stop that worked. An operator calling the unload endpoint got HTTP 200 while +// the backend kept running and holding its VRAM. +type BackendStopReply struct { + Success bool `json:"success"` + Error string `json:"error,omitempty"` + + // StoppedProcessKeys names every `modelID#replica` process the worker + // terminated while serving this request. + StoppedProcessKeys []string `json:"stopped_process_keys,omitempty"` + + // ReportsStoppedProcesses distinguishes "this worker enumerates what it + // stopped and stopped nothing" from "this worker predates the field", the + // same way BackendDeleteReply does. Both send an empty list and only the + // first is authoritative, so a controller that cannot tell them apart would + // read silence as a completed stop, the exact conclusion this reply exists + // to prevent. + ReportsStoppedProcesses bool `json:"reports_stopped_processes,omitempty"` +} + +// BackendDeleteRequest is the payload for a backend.delete control request. +type BackendDeleteRequest struct { + Backend string `json:"backend"` +} + +// BackendDeleteReply is the response from a backend.delete control request. +type BackendDeleteReply struct { + Success bool `json:"success"` + Error string `json:"error,omitempty"` + + // StoppedProcessKeys names every `modelID#replica` process the worker + // terminated while serving this delete. Stopping a process returns its gRPC + // port to the worker's allocator, so any NodeModel row still pointing at + // that address becomes a live misroute the moment an unrelated backend + // binds the recycled port: probeHealth verifies liveness, not identity, so + // the request is served by the wrong backend rather than failing. The + // controller uses these keys to drop the rows eagerly. + StoppedProcessKeys []string `json:"stopped_process_keys,omitempty"` + + // ReportsStoppedProcesses distinguishes "this worker enumerates what it + // stopped and stopped nothing" from "this worker predates the field". Both + // send an empty list, and only the first is authoritative. Without this + // flag a controller cannot tell them apart and would eventually be tempted + // to read silence as a completed cleanup, which is precisely the wrong + // conclusion against an older worker. + ReportsStoppedProcesses bool `json:"reports_stopped_processes,omitempty"` +} diff --git a/core/services/workerctl/backend_test.go b/core/services/workerctl/backend_test.go new file mode 100644 index 000000000..045aedf6f --- /dev/null +++ b/core/services/workerctl/backend_test.go @@ -0,0 +1,20 @@ +package workerctl_test + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/workerctl" +) + +var _ = Describe("BackendUpgradeRequest", func() { + It("carries backend name, galleries JSON, and replica index", func() { + req := workerctl.BackendUpgradeRequest{ + Backend: "llama-cpp", + BackendGalleries: `[{"name":"x"}]`, + ReplicaIndex: 2, + } + Expect(req.Backend).To(Equal("llama-cpp")) + Expect(req.ReplicaIndex).To(BeEquivalentTo(2)) + }) +}) diff --git a/core/services/workerctl/doc.go b/core/services/workerctl/doc.go new file mode 100644 index 000000000..f6f80ddde --- /dev/null +++ b/core/services/workerctl/doc.go @@ -0,0 +1,9 @@ +// Package workerctl holds the request and reply payloads of the worker control +// verbs (backend install, upgrade, list, stop and delete, model stop, unload +// and delete, running models, and the file staging verbs). +// +// The payloads live apart from any carrier so that every transport that serves +// or sends a verb decodes the same structs. The package imports only the +// standard library, which keeps it a leaf that both the controller and the +// worker can depend on without an import cycle. +package workerctl diff --git a/core/services/workerctl/files.go b/core/services/workerctl/files.go new file mode 100644 index 000000000..4bd8fea60 --- /dev/null +++ b/core/services/workerctl/files.go @@ -0,0 +1,64 @@ +package workerctl + +// The file staging verbs let the controller move model and request files +// between shared object storage and a worker's local cache. The success +// fields carry omitempty so a reply holds only the key that the outcome sets, +// which matches the single-key replies that workers already send. + +// FileEnsureRequest asks a worker to download an object storage key into its +// local cache. +type FileEnsureRequest struct { + Key string `json:"key"` +} + +// FileEnsureReply carries the local path of the cached file, or an error. +type FileEnsureReply struct { + LocalPath string `json:"local_path,omitempty"` + Error string `json:"error,omitempty"` +} + +// FileStageRequest asks a worker to upload a local file to object storage +// under Key. +type FileStageRequest struct { + LocalPath string `json:"local_path"` + Key string `json:"key"` +} + +// FileStageReply carries the key the file was uploaded under, or an error. +type FileStageReply struct { + Key string `json:"key,omitempty"` + Error string `json:"error,omitempty"` +} + +// FileTempRequest asks a worker to allocate a temporary file. +type FileTempRequest struct{} + +// FileTempReply carries the path of the allocated temporary file, or an error. +type FileTempReply struct { + LocalPath string `json:"local_path,omitempty"` + Error string `json:"error,omitempty"` +} + +// FileListDirRequest asks a worker to list the files under a key prefix. +type FileListDirRequest struct { + KeyPrefix string `json:"key_prefix"` +} + +// FileListDirReply carries the listed files, or an error. +type FileListDirReply struct { + Files []string `json:"files,omitempty"` + Error string `json:"error,omitempty"` +} + +// FileReleaseRequest asks a worker to evict ephemeral cache entries. Key names +// one exact key; RequestID names every key staged for one inference. A worker +// takes the RequestID form whenever RequestID is set, so a sender fills one. +type FileReleaseRequest struct { + Key string `json:"key,omitempty"` + RequestID string `json:"request_id,omitempty"` +} + +// FileReleaseReply carries an error, or nothing on success. +type FileReleaseReply struct { + Error string `json:"error,omitempty"` +} diff --git a/core/services/workerctl/model.go b/core/services/workerctl/model.go new file mode 100644 index 000000000..b9074bb63 --- /dev/null +++ b/core/services/workerctl/model.go @@ -0,0 +1,59 @@ +package workerctl + +type ModelStopRequest struct { + ModelName string `json:"model_name"` + ProcessKey string `json:"process_key"` + ExpectedAddress string `json:"expected_address"` + Force bool `json:"force,omitempty"` + ConfigRevision string `json:"config_revision,omitempty"` +} + +type ModelStopReply struct { + Matched bool `json:"matched"` + Freed bool `json:"freed"` + Terminated bool `json:"terminated"` + ProcessKey string `json:"process_key"` + Address string `json:"address,omitempty"` + Error string `json:"error,omitempty"` +} + +// ModelUnloadRequest is the payload for a model.unload control request. +type ModelUnloadRequest struct { + ModelName string `json:"model_name"` + Address string `json:"address,omitempty"` // gRPC address of the backend process to unload from +} + +// ModelUnloadReply is the response from a model.unload control request. +type ModelUnloadReply struct { + Success bool `json:"success"` + Error string `json:"error,omitempty"` +} + +// ModelDeleteRequest is the payload for a model.delete control request. +type ModelDeleteRequest struct { + ModelName string `json:"model_name"` +} + +// ModelDeleteReply is the response from a model.delete control request. +type ModelDeleteReply struct { + Success bool `json:"success"` + Error string `json:"error,omitempty"` +} + +// ModelsRunningRequest is the payload for a models.running control request. +type ModelsRunningRequest struct{} + +// ModelsRunningReply is the response from a models.running control request. +type ModelsRunningReply struct { + Models []RunningModelInfo `json:"models"` + Error string `json:"error,omitempty"` +} + +// RunningModelInfo identifies one live backend process on a worker. The triple +// is isomorphic to a controller NodeModel row's (model_name, replica_index, +// address), which is what lets the reconciler diff the two directly. +type RunningModelInfo struct { + ModelID string `json:"model_id"` + ReplicaIndex int `json:"replica_index"` + Address string `json:"address,omitempty"` +} diff --git a/core/services/workerctl/progress.go b/core/services/workerctl/progress.go new file mode 100644 index 000000000..ecdfbb46e --- /dev/null +++ b/core/services/workerctl/progress.go @@ -0,0 +1,29 @@ +package workerctl + +// Phase values published on the BackendInstallProgressEvent.Phase field. +// Defined as exported constants so producer (worker install handler) and +// consumer (master bridge into OpStatus) share a single source of truth +// instead of two copies of the literal string. +const ( + PhaseResolving = "resolving" // worker is locating the gallery / image manifest + PhaseDownloading = "downloading" // worker is actively pulling layers + PhaseExtracting = "extracting" // worker is unpacking the downloaded archive + PhaseStarting = "starting" // worker is spawning the gRPC backend process +) + +// BackendInstallProgressEvent is the wire payload published by a worker to +// nodes..backend.install..progress while a long-running install +// is in flight. Transient: dropped events are acceptable, the master relies +// on BackendInstallReply for ground truth on success/failure. +// +// Phase holds one of the Phase* constants above. +type BackendInstallProgressEvent struct { + OpID string `json:"op_id"` + NodeID string `json:"node_id"` + Backend string `json:"backend"` + FileName string `json:"file_name,omitempty"` + Current string `json:"current,omitempty"` // human-readable size, e.g. "412 MB" + Total string `json:"total,omitempty"` // human-readable size, e.g. "2.1 GB" + Percentage float64 `json:"percentage"` + Phase string `json:"phase,omitempty"` +} diff --git a/core/services/workerctl/progress_test.go b/core/services/workerctl/progress_test.go new file mode 100644 index 000000000..d81f19eee --- /dev/null +++ b/core/services/workerctl/progress_test.go @@ -0,0 +1,48 @@ +package workerctl_test + +import ( + "encoding/json" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/workerctl" +) + +var _ = Describe("Phase constants", func() { + // Pin the wire-format string values. A future refactor that renames + // a constant must NOT silently change the JSON value the master + // receives or break consumers that switch on Phase. + DescribeTable("phase constant", + func(actual, expected string) { + Expect(actual).To(Equal(expected)) + }, + Entry("resolving", workerctl.PhaseResolving, "resolving"), + Entry("downloading", workerctl.PhaseDownloading, "downloading"), + Entry("extracting", workerctl.PhaseExtracting, "extracting"), + Entry("starting", workerctl.PhaseStarting, "starting"), + ) +}) + +var _ = Describe("BackendInstallProgress", func() { + Context("BackendInstallProgressEvent", func() { + It("JSON round-trips with all known fields", func() { + ev := workerctl.BackendInstallProgressEvent{ + OpID: "op-123", + NodeID: "node-abc", + Backend: "vllm", + FileName: "vllm-cpu.tar.zst", + Current: "412 MB", + Total: "2.1 GB", + Percentage: 19.6, + Phase: "downloading", + } + raw, err := json.Marshal(ev) + Expect(err).ToNot(HaveOccurred()) + + var got workerctl.BackendInstallProgressEvent + Expect(json.Unmarshal(raw, &got)).To(Succeed()) + Expect(got).To(Equal(ev)) + }) + }) +}) diff --git a/core/services/workerctl/wire_test.go b/core/services/workerctl/wire_test.go new file mode 100644 index 000000000..319525686 --- /dev/null +++ b/core/services/workerctl/wire_test.go @@ -0,0 +1,153 @@ +package workerctl_test + +import ( + "encoding/json" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/workerctl" +) + +// Workers and controllers of different versions talk to each other during a +// rolling update, so the JSON field names and omitempty choices are a wire +// contract. These specs pin the exact bytes rather than a round trip, because +// a round trip still passes after a tag is renamed on both sides at once. +var _ = Describe("Control payload wire format", func() { + DescribeTable("marshals to the pinned bytes", + func(v any, expected string) { + raw, err := json.Marshal(v) + Expect(err).ToNot(HaveOccurred()) + Expect(string(raw)).To(Equal(expected)) + }, + Entry("BackendInstallRequest, every field set", workerctl.BackendInstallRequest{ + Backend: "b", ModelID: "m", BackendGalleries: "g", URI: "u", Name: "n", Alias: "a", + ReplicaIndex: 2, Force: true, OpID: "op", + }, `{"backend":"b","model_id":"m","backend_galleries":"g","uri":"u","name":"n","alias":"a","replica_index":2,"force":true,"op_id":"op"}`), + Entry("BackendInstallRequest, zero value", workerctl.BackendInstallRequest{}, `{"backend":""}`), + + Entry("BackendInstallReply, every field set", workerctl.BackendInstallReply{ + Success: true, Address: "h:1", Error: "e", + }, `{"success":true,"address":"h:1","error":"e"}`), + Entry("BackendInstallReply, zero value", workerctl.BackendInstallReply{}, `{"success":false}`), + + Entry("BackendUpgradeRequest, every field set", workerctl.BackendUpgradeRequest{ + Backend: "b", BackendGalleries: "g", URI: "u", Name: "n", Alias: "a", ReplicaIndex: 2, OpID: "op", + }, `{"backend":"b","backend_galleries":"g","uri":"u","name":"n","alias":"a","replica_index":2,"op_id":"op"}`), + Entry("BackendUpgradeRequest, zero value", workerctl.BackendUpgradeRequest{}, `{"backend":""}`), + + Entry("BackendUpgradeReply, every field set", workerctl.BackendUpgradeReply{ + Success: true, Error: "e", StoppedProcessKeys: []string{"k#0"}, ReportsStoppedProcesses: true, + }, `{"success":true,"error":"e","stopped_process_keys":["k#0"],"reports_stopped_processes":true}`), + Entry("BackendUpgradeReply, zero value", workerctl.BackendUpgradeReply{}, `{"success":false}`), + + Entry("BackendListRequest", workerctl.BackendListRequest{}, `{}`), + + Entry("BackendListReply, every field set", workerctl.BackendListReply{ + Backends: []workerctl.NodeBackendInfo{{ + Name: "n", IsSystem: true, IsMeta: true, InstalledAt: "t", GalleryURL: "gu", + Version: "v", URI: "u", Digest: "d", + }}, + Error: "e", + }, `{"backends":[{"name":"n","is_system":true,"is_meta":true,"installed_at":"t","gallery_url":"gu","version":"v","uri":"u","digest":"d"}],"error":"e"}`), + Entry("BackendListReply, zero value", workerctl.BackendListReply{}, `{"backends":null}`), + + Entry("NodeBackendInfo, every field set", workerctl.NodeBackendInfo{ + Name: "n", IsSystem: true, IsMeta: true, InstalledAt: "t", GalleryURL: "gu", + Version: "v", URI: "u", Digest: "d", + }, `{"name":"n","is_system":true,"is_meta":true,"installed_at":"t","gallery_url":"gu","version":"v","uri":"u","digest":"d"}`), + Entry("NodeBackendInfo, zero value", workerctl.NodeBackendInfo{}, `{"name":"","is_system":false,"is_meta":false}`), + + Entry("BackendStopRequest, every field set", workerctl.BackendStopRequest{Backend: "b", Force: true}, + `{"backend":"b","force":true}`), + Entry("BackendStopRequest, zero value", workerctl.BackendStopRequest{}, `{"backend":""}`), + + Entry("BackendStopReply, every field set", workerctl.BackendStopReply{ + Success: true, Error: "e", StoppedProcessKeys: []string{"k#0"}, ReportsStoppedProcesses: true, + }, `{"success":true,"error":"e","stopped_process_keys":["k#0"],"reports_stopped_processes":true}`), + Entry("BackendStopReply, zero value", workerctl.BackendStopReply{}, `{"success":false}`), + + Entry("ModelStopRequest, every field set", workerctl.ModelStopRequest{ + ModelName: "m", ProcessKey: "k", ExpectedAddress: "a", Force: true, ConfigRevision: "r", + }, `{"model_name":"m","process_key":"k","expected_address":"a","force":true,"config_revision":"r"}`), + Entry("ModelStopRequest, zero value", workerctl.ModelStopRequest{}, + `{"model_name":"","process_key":"","expected_address":""}`), + + Entry("ModelStopReply, every field set", workerctl.ModelStopReply{ + Matched: true, Freed: true, Terminated: true, ProcessKey: "k", Address: "a", Error: "e", + }, `{"matched":true,"freed":true,"terminated":true,"process_key":"k","address":"a","error":"e"}`), + Entry("ModelStopReply, zero value", workerctl.ModelStopReply{}, + `{"matched":false,"freed":false,"terminated":false,"process_key":""}`), + + Entry("BackendDeleteRequest", workerctl.BackendDeleteRequest{Backend: "b"}, `{"backend":"b"}`), + + Entry("BackendDeleteReply, every field set", workerctl.BackendDeleteReply{ + Success: true, Error: "e", StoppedProcessKeys: []string{"k#0"}, ReportsStoppedProcesses: true, + }, `{"success":true,"error":"e","stopped_process_keys":["k#0"],"reports_stopped_processes":true}`), + Entry("BackendDeleteReply, zero value", workerctl.BackendDeleteReply{}, `{"success":false}`), + + Entry("ModelUnloadRequest, every field set", workerctl.ModelUnloadRequest{ModelName: "m", Address: "a"}, + `{"model_name":"m","address":"a"}`), + Entry("ModelUnloadRequest, zero value", workerctl.ModelUnloadRequest{}, `{"model_name":""}`), + + Entry("ModelUnloadReply, every field set", workerctl.ModelUnloadReply{Success: true, Error: "e"}, + `{"success":true,"error":"e"}`), + Entry("ModelUnloadReply, zero value", workerctl.ModelUnloadReply{}, `{"success":false}`), + + Entry("ModelDeleteRequest", workerctl.ModelDeleteRequest{ModelName: "m"}, `{"model_name":"m"}`), + + Entry("ModelDeleteReply, every field set", workerctl.ModelDeleteReply{Success: true, Error: "e"}, + `{"success":true,"error":"e"}`), + Entry("ModelDeleteReply, zero value", workerctl.ModelDeleteReply{}, `{"success":false}`), + + Entry("ModelsRunningRequest", workerctl.ModelsRunningRequest{}, `{}`), + + Entry("ModelsRunningReply, every field set", workerctl.ModelsRunningReply{ + Models: []workerctl.RunningModelInfo{{ModelID: "m", ReplicaIndex: 1, Address: "a"}}, + Error: "e", + }, `{"models":[{"model_id":"m","replica_index":1,"address":"a"}],"error":"e"}`), + Entry("ModelsRunningReply, zero value", workerctl.ModelsRunningReply{}, `{"models":null}`), + + Entry("RunningModelInfo, every field set", workerctl.RunningModelInfo{ModelID: "m", ReplicaIndex: 1, Address: "a"}, + `{"model_id":"m","replica_index":1,"address":"a"}`), + Entry("RunningModelInfo, zero value", workerctl.RunningModelInfo{}, `{"model_id":"","replica_index":0}`), + + Entry("BackendInstallProgressEvent, every field set", workerctl.BackendInstallProgressEvent{ + OpID: "op", NodeID: "n", Backend: "b", FileName: "f", Current: "1 MB", Total: "2 MB", + Percentage: 19.6, Phase: workerctl.PhaseDownloading, + }, `{"op_id":"op","node_id":"n","backend":"b","file_name":"f","current":"1 MB","total":"2 MB","percentage":19.6,"phase":"downloading"}`), + Entry("BackendInstallProgressEvent, zero value", workerctl.BackendInstallProgressEvent{}, + `{"op_id":"","node_id":"","backend":"","percentage":0}`), + ) + + // The file staging replies below must equal the single-key maps that the + // worker's file staging handlers marshal today, so a worker that switches to + // these structs sends the same bytes. + DescribeTable("file staging payloads marshal to the pinned bytes", + func(v any, expected string) { + raw, err := json.Marshal(v) + Expect(err).ToNot(HaveOccurred()) + Expect(string(raw)).To(Equal(expected)) + }, + Entry("FileEnsureRequest", workerctl.FileEnsureRequest{Key: "k"}, `{"key":"k"}`), + Entry("FileEnsureReply, success", workerctl.FileEnsureReply{LocalPath: "x"}, `{"local_path":"x"}`), + Entry("FileEnsureReply, error", workerctl.FileEnsureReply{Error: "e"}, `{"error":"e"}`), + + Entry("FileStageRequest", workerctl.FileStageRequest{LocalPath: "p", Key: "k"}, `{"local_path":"p","key":"k"}`), + Entry("FileStageReply, success", workerctl.FileStageReply{Key: "x"}, `{"key":"x"}`), + Entry("FileStageReply, error", workerctl.FileStageReply{Error: "e"}, `{"error":"e"}`), + + Entry("FileTempRequest", workerctl.FileTempRequest{}, `{}`), + Entry("FileTempReply, success", workerctl.FileTempReply{LocalPath: "x"}, `{"local_path":"x"}`), + Entry("FileTempReply, error", workerctl.FileTempReply{Error: "e"}, `{"error":"e"}`), + + Entry("FileListDirRequest", workerctl.FileListDirRequest{KeyPrefix: "p/"}, `{"key_prefix":"p/"}`), + Entry("FileListDirReply, success", workerctl.FileListDirReply{Files: []string{"a"}}, `{"files":["a"]}`), + Entry("FileListDirReply, error", workerctl.FileListDirReply{Error: "e"}, `{"error":"e"}`), + + Entry("FileReleaseRequest, exact key", workerctl.FileReleaseRequest{Key: "k"}, `{"key":"k"}`), + Entry("FileReleaseRequest, request id", workerctl.FileReleaseRequest{RequestID: "r"}, `{"request_id":"r"}`), + Entry("FileReleaseReply, success", workerctl.FileReleaseReply{}, `{}`), + Entry("FileReleaseReply, error", workerctl.FileReleaseReply{Error: "e"}, `{"error":"e"}`), + ) +}) diff --git a/core/services/workerctl/workerctl_suite_test.go b/core/services/workerctl/workerctl_suite_test.go new file mode 100644 index 000000000..a47f1c8ff --- /dev/null +++ b/core/services/workerctl/workerctl_suite_test.go @@ -0,0 +1,13 @@ +package workerctl_test + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestWorkerctl(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Workerctl test suite") +} diff --git a/docs/content/features/distributed-mode.md b/docs/content/features/distributed-mode.md index 07d1d7378..96f6fd08f 100644 --- a/docs/content/features/distributed-mode.md +++ b/docs/content/features/distributed-mode.md @@ -852,6 +852,8 @@ Agent workers: - Handle MCP tool discovery and execution requests from the frontend - Get auto-provisioned API keys during registration for calling the inference API +`LOCALAI_AGENT_SUBJECT` (default `agent.execute`) must be a subject that LocalAI serves. Use the `agent` root, for example `agent.execute`. The worker refuses to start with a subject whose root LocalAI does not serve (for example `tenant-a.agent.execute`) or with a `>` wildcard, because no message is carried on those subjects. + In the docker-compose setup, the agent worker mounts the Docker socket so it can run MCP stdio servers (e.g., `docker run` commands): ```yaml diff --git a/tests/e2e/distributed/agent_native_executor_test.go b/tests/e2e/distributed/agent_native_executor_test.go index b58e32141..921e94113 100644 --- a/tests/e2e/distributed/agent_native_executor_test.go +++ b/tests/e2e/distributed/agent_native_executor_test.go @@ -314,7 +314,7 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f }) Context("NATSDispatcher", func() { - It("should dispatch chat via NATS and receive response", func() { + It("should run a chat enqueued on the agent-run queue", func() { bridge := agents.NewEventBridge(infra.NC, nil, "test-instance") configs := &mockConfigProvider{configs: map[string]*agents.AgentConfig{ @@ -340,40 +340,38 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f Expect(err).ToNot(HaveOccurred()) defer sub.Unsubscribe() - adapter := infra.NC - dispatcher := agents.NewNATSDispatcher(adapter, bridge, configs, "http://localhost:8080", "test-key", "agent.test.execute", "test-workers", 0) + // The consumer and the enqueue both use the default agent-run route. + // Nothing listens on the API URL, so the run fails after the worker + // has taken it, which is what the error message below proves. + dispatcher := agents.NewNATSDispatcher(messaging.NewNATSWorkConsumer(infra.NC), bridge, configs, "http://127.0.0.1:1", "test-key", 0) + Expect(dispatcher.Start(infra.Ctx)).To(Succeed()) + defer func() { _ = dispatcher.Stop() }() + FlushNATS(infra.NC) - err = dispatcher.Start(infra.Ctx) - Expect(err).ToNot(HaveOccurred()) + // No Config in the payload: the worker resolves it through ConfigProvider. + Expect(messaging.NewNATSWorkQueue(infra.NC).Enqueue(infra.Ctx, messaging.WorkAgentRun, agents.AgentChatEvent{ + AgentName: "test-agent", + UserID: "user1", + Message: "Hello", + MessageID: "msg-queued-001", + Role: agents.RoleUser, + })).To(Succeed()) - // Dispatch a chat - messageID, err := dispatcher.Dispatch("user1", "test-agent", "Hello") - Expect(err).ToNot(HaveOccurred()) - Expect(messageID).ToNot(BeEmpty()) - - // Wait for events (user message + processing status should arrive immediately) - Eventually(func() int { - eventMu.Lock() - defer eventMu.Unlock() - return len(receivedEvents) - }, "5s").Should(BeNumerically(">=", 2)) - - // Verify user message was published - eventMu.Lock() - hasUserMsg := false - hasProcessing := false - for _, evt := range receivedEvents { - if evt.EventType == "json_message" && evt.Sender == "user" { - hasUserMsg = true - } - if evt.EventType == "json_message_status" { - hasProcessing = true + hasEvent := func(eventType, sender, messageID string) func() bool { + return func() bool { + eventMu.Lock() + defer eventMu.Unlock() + for _, evt := range receivedEvents { + if evt.EventType == eventType && evt.Sender == sender && (messageID == "" || evt.MessageID == messageID) { + return true + } + } + return false } } - eventMu.Unlock() - - Expect(hasUserMsg).To(BeTrue(), "user message should be published immediately") - Expect(hasProcessing).To(BeTrue(), "processing status should be published") + Eventually(hasEvent("json_message_status", "", ""), "10s").Should(BeTrue(), "the worker should report processing") + Eventually(hasEvent("json_message", agents.RoleAgent, "msg-queued-001-error"), "30s").Should(BeTrue(), + "the worker should run the enqueued message and report its failure") }) It("should handle cancellation via EventBridge", func() { @@ -402,7 +400,7 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f // Create dispatcher with NO ConfigProvider (simulating DB-free worker) adapter := infra.NC - dispatcher := agents.NewNATSDispatcher(adapter, bridge, nil, "http://localhost:8080", "test-key", "agent.enriched.execute", "enriched-workers", 0) + dispatcher := agents.NewNATSDispatcher(messaging.NewNATSWorkConsumer(adapter, messaging.WithAgentRunRoute("agent.enriched.execute", "enriched-workers")), bridge, nil, "http://localhost:8080", "test-key", 0) Expect(dispatcher.Start(infra.Ctx)).To(Succeed()) // Subscribe to events to verify processing @@ -625,7 +623,7 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f }) Context("Full Distributed Chat Flow", func() { - It("should dispatch chat via NATS, execute, and publish response via EventBridge", func() { + It("should enqueue a stored agent chat, run it on the worker, and publish events via EventBridge", func() { bridge := agents.NewEventBridge(infra.NC, nil, "flow-test") // Store agent config in PostgreSQL @@ -662,34 +660,51 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f "flow-agent": &cfg, }} - dispatcher := agents.NewNATSDispatcher(adapter, bridge, configs, "http://localhost:8080", "test-key", "agent.flow.execute", "flow-workers", 0) + // Nothing listens on the API URL, so the run fails after the worker + // has taken it; the error message below is the proof it did. + dispatcher := agents.NewNATSDispatcher(messaging.NewNATSWorkConsumer(adapter), bridge, configs, "http://127.0.0.1:1", "test-key", 0) Expect(dispatcher.Start(infra.Ctx)).To(Succeed()) + defer func() { _ = dispatcher.Stop() }() + FlushNATS(infra.NC) - // Dispatch - messageID, err := dispatcher.Dispatch("user1", "flow-agent", "Hello flow test") + // Act as the frontend does for a chat: show the user message, then + // enqueue the run with the stored config embedded. + const messageID = "msg-flow-001" + Expect(bridge.PublishMessage("flow-agent", "user1", agents.RoleUser, "Hello flow test", messageID+"-user")).To(Succeed()) + rec, err := store.GetConfig("user1", "flow-agent") Expect(err).ToNot(HaveOccurred()) - Expect(messageID).ToNot(BeEmpty()) + var stored agents.AgentConfig + Expect(agents.ParseConfigJSON(rec.ConfigJSON, &stored)).To(Succeed()) + Expect(messaging.NewNATSWorkQueue(infra.NC).Enqueue(infra.Ctx, messaging.WorkAgentRun, agents.AgentChatEvent{ + AgentName: "flow-agent", + UserID: "user1", + Message: "Hello flow test", + MessageID: messageID, + Role: agents.RoleUser, + Config: &stored, + })).To(Succeed()) - // User message + processing status should arrive immediately - Eventually(func() int { - eventMu.Lock() - defer eventMu.Unlock() - return len(receivedEvents) - }, "5s").Should(BeNumerically(">=", 2)) - - eventMu.Lock() - var hasUser, hasProcessing bool - for _, evt := range receivedEvents { - if evt.EventType == "json_message" && evt.Sender == "user" && evt.Content == "Hello flow test" { - hasUser = true - } - if evt.EventType == "json_message_status" { - hasProcessing = true + hasEvent := func(match func(agents.AgentEvent) bool) func() bool { + return func() bool { + eventMu.Lock() + defer eventMu.Unlock() + for _, evt := range receivedEvents { + if match(evt) { + return true + } + } + return false } } - eventMu.Unlock() - Expect(hasUser).To(BeTrue(), "expected user message event") - Expect(hasProcessing).To(BeTrue(), "expected processing status event") + Eventually(hasEvent(func(evt agents.AgentEvent) bool { + return evt.EventType == "json_message" && evt.Sender == agents.RoleUser && evt.Content == "Hello flow test" + }), "5s").Should(BeTrue(), "expected user message event") + Eventually(hasEvent(func(evt agents.AgentEvent) bool { + return evt.EventType == "json_message_status" + }), "10s").Should(BeTrue(), "expected processing status event from the worker") + Eventually(hasEvent(func(evt agents.AgentEvent) bool { + return evt.EventType == "json_message" && evt.Sender == agents.RoleAgent && evt.MessageID == messageID+"-error" + }), "30s").Should(BeTrue(), "expected the worker to run the enqueued message") }) }) @@ -758,7 +773,7 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f Expect(err).ToNot(HaveOccurred()) defer sub.Unsubscribe() - dispatcher := agents.NewNATSDispatcher(adapter, bridge, configs, "http://localhost:8080", "test-key", "agent.bg.execute", "bg-workers", 0) + dispatcher := agents.NewNATSDispatcher(messaging.NewNATSWorkConsumer(adapter, messaging.WithAgentRunRoute("agent.bg.execute", "bg-workers")), bridge, configs, "http://localhost:8080", "test-key", 0) Expect(dispatcher.Start(infra.Ctx)).To(Succeed()) // Dispatch as background/system role @@ -866,7 +881,9 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f // Subscribe to NATS to capture background run events var receivedEvents []agents.AgentChatEvent var eventMu sync.Mutex - sub, err := infra.NC.Subscribe("agent.sched.execute", func(data []byte) { + // A plain subscription sees every publish, alongside any queue + // group that also listens on the agent-run subject. + sub, err := infra.NC.Subscribe(messaging.SubjectAgentExecute, func(data []byte) { var evt agents.AgentChatEvent if json.Unmarshal(data, &evt) == nil { eventMu.Lock() @@ -878,8 +895,7 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f defer sub.Unsubscribe() // Start scheduler with short poll interval for testing - adapter := infra.NC - scheduler := agents.NewAgentScheduler(db, adapter, store, "agent.sched.execute") + scheduler := agents.NewAgentScheduler(db, messaging.NewNATSWorkQueue(infra.NC), store) schedCtx, schedCancel := context.WithCancel(infra.Ctx) defer schedCancel() @@ -1081,7 +1097,7 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f adapter := infra.NC // Point dispatcher at our mock LLM server - dispatcher := agents.NewNATSDispatcher(adapter, bridge, nil, llmURL, "test-key", "agent.e2e.execute", "e2e-workers", 0) + dispatcher := agents.NewNATSDispatcher(messaging.NewNATSWorkConsumer(adapter, messaging.WithAgentRunRoute("agent.e2e.execute", "e2e-workers")), bridge, nil, llmURL, "test-key", 0) Expect(dispatcher.Start(infra.Ctx)).To(Succeed()) FlushNATS(infra.NC) @@ -1156,7 +1172,7 @@ var _ = Describe("Native Agent Executor", Label("Distributed", "AgentNative"), f defer sub.Unsubscribe() adapter := infra.NC - dispatcher := agents.NewNATSDispatcher(adapter, bridge, nil, llmURL, "test-key", "agent.bg-e2e.execute", "bg-e2e-workers", 0) + dispatcher := agents.NewNATSDispatcher(messaging.NewNATSWorkConsumer(adapter, messaging.WithAgentRunRoute("agent.bg-e2e.execute", "bg-e2e-workers")), bridge, nil, llmURL, "test-key", 0) Expect(dispatcher.Start(infra.Ctx)).To(Succeed()) FlushNATS(infra.NC) diff --git a/tests/e2e/distributed/backend_logs_test.go b/tests/e2e/distributed/backend_logs_test.go index 79dea3902..e721f58cf 100644 --- a/tests/e2e/distributed/backend_logs_test.go +++ b/tests/e2e/distributed/backend_logs_test.go @@ -9,6 +9,7 @@ import ( "net/http/httptest" "net/url" "os" + "sync" "time" "github.com/gorilla/websocket" @@ -343,7 +344,7 @@ var _ = Describe("Distributed Backend Log Streaming", Label("Distributed"), func // Create an Echo test server with the proxy endpoint e := echo.New() - e.GET("/api/nodes/:id/backend-logs", localai.NodeBackendLogsListEndpoint(registry, token)) + e.GET("/api/nodes/:id/backend-logs", localai.NodeBackendLogsListEndpoint(registry, token, nodes.DirectWorkerNetDialer())) req := httptest.NewRequest("GET", fmt.Sprintf("/api/nodes/%s/backend-logs", node.ID), nil) rec := httptest.NewRecorder() @@ -365,7 +366,7 @@ var _ = Describe("Distributed Backend Log Streaming", Label("Distributed"), func Expect(registry.Register(context.Background(), node, true)).To(Succeed()) e := echo.New() - e.GET("/api/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsLinesEndpoint(registry, token)) + e.GET("/api/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsLinesEndpoint(registry, token, nodes.DirectWorkerNetDialer())) req := httptest.NewRequest("GET", fmt.Sprintf("/api/nodes/%s/backend-logs/remote-model", node.ID), nil) rec := httptest.NewRecorder() @@ -382,7 +383,7 @@ var _ = Describe("Distributed Backend Log Streaming", Label("Distributed"), func It("should return 404 for unknown node ID", func() { e := echo.New() - e.GET("/api/nodes/:id/backend-logs", localai.NodeBackendLogsListEndpoint(registry, token)) + e.GET("/api/nodes/:id/backend-logs", localai.NodeBackendLogsListEndpoint(registry, token, nodes.DirectWorkerNetDialer())) req := httptest.NewRequest("GET", "/api/nodes/nonexistent-id/backend-logs", nil) rec := httptest.NewRecorder() @@ -426,7 +427,7 @@ var _ = Describe("Distributed Backend Log Streaming", Label("Distributed"), func // Start Echo server with the WebSocket proxy route e := echo.New() - e.GET("/ws/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsWSEndpoint(registry, token)) + e.GET("/ws/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsWSEndpoint(registry, token, nodes.DirectWorkerNetDialer())) lis, err := net.Listen("tcp", "127.0.0.1:0") Expect(err).ToNot(HaveOccurred()) @@ -503,6 +504,125 @@ var _ = Describe("Distributed Backend Log Streaming", Label("Distributed"), func } }) }) + + Context("Frontend proxy through the per-node worker dialer", func() { + var ( + infra *TestInfra + registry *nodes.NodeRegistry + logStore *model.BackendLogStore + workerAddr string + workerClean func() + token string + echoServer *http.Server + echoAddr string + dialedMu sync.Mutex + dialed []string + ) + + BeforeEach(func() { + infra = SetupInfra("localai_backend_logs_dialer_test") + + db, err := gorm.Open(pgdriver.Open(infra.PGURL), &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + }) + Expect(err).ToNot(HaveOccurred()) + + registry, err = nodes.NewNodeRegistry(db) + Expect(err).ToNot(HaveOccurred()) + + token = "dialer-proxy-token" + logStore = model.NewBackendLogStore(1000) + logStore.AppendLine("dialed-model", "stdout", "line through the dialer") + + workerAddr, workerClean, err = startTestFileTransferServerWithLogs(token, logStore) + Expect(err).ToNot(HaveOccurred()) + + dialedMu.Lock() + dialed = nil + dialedMu.Unlock() + + // The node's advertised address is unresolvable on purpose: only a + // proxy that dials through the per-node dialer can reach the worker. + var d net.Dialer + dialFor := func(nodeID string) func(context.Context, string, string) (net.Conn, error) { + return func(ctx context.Context, network, _ string) (net.Conn, error) { + dialedMu.Lock() + dialed = append(dialed, nodeID) + dialedMu.Unlock() + return d.DialContext(ctx, network, workerAddr) + } + } + + e := echo.New() + e.GET("/api/nodes/:id/backend-logs", localai.NodeBackendLogsListEndpoint(registry, token, dialFor)) + e.GET("/api/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsLinesEndpoint(registry, token, dialFor)) + e.GET("/ws/nodes/:id/backend-logs/:modelId", localai.NodeBackendLogsWSEndpoint(registry, token, dialFor)) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + Expect(err).ToNot(HaveOccurred()) + echoAddr = lis.Addr().String() + echoServer = &http.Server{Handler: e} + go func() { _ = echoServer.Serve(lis) }() + + Expect(registry.Register(context.Background(), &nodes.BackendNode{ + ID: "n1", + Name: "dialer-node", + Address: "127.0.0.1:50051", + HTTPAddress: "n1.worker.invalid:80", + }, true)).To(Succeed()) + }) + + AfterEach(func() { + if echoServer != nil { + _ = echoServer.Close() + } + if workerClean != nil { + workerClean() + } + }) + + dialedNodes := func() []string { + dialedMu.Lock() + defer dialedMu.Unlock() + return append([]string(nil), dialed...) + } + + It("routes backend logs list, lines and WebSocket through the dialer", func() { + client := &http.Client{Timeout: 10 * time.Second} + + resp, err := client.Get(fmt.Sprintf("http://%s/api/nodes/n1/backend-logs", echoAddr)) + Expect(err).ToNot(HaveOccurred()) + var models []string + Expect(resp.StatusCode).To(Equal(http.StatusOK)) + Expect(json.NewDecoder(resp.Body).Decode(&models)).To(Succeed()) + Expect(resp.Body.Close()).To(Succeed()) + Expect(models).To(ContainElement("dialed-model")) + Expect(dialedNodes()).To(Equal([]string{"n1"})) + + resp, err = client.Get(fmt.Sprintf("http://%s/api/nodes/n1/backend-logs/dialed-model", echoAddr)) + Expect(err).ToNot(HaveOccurred()) + var lines []model.BackendLogLine + Expect(resp.StatusCode).To(Equal(http.StatusOK)) + Expect(json.NewDecoder(resp.Body).Decode(&lines)).To(Succeed()) + Expect(resp.Body.Close()).To(Succeed()) + Expect(lines).To(HaveLen(1)) + Expect(lines[0].Text).To(Equal("line through the dialer")) + Expect(dialedNodes()).To(Equal([]string{"n1", "n1"})) + + wsDialer := websocket.Dialer{HandshakeTimeout: 5 * time.Second} + conn, _, err := wsDialer.Dial(fmt.Sprintf("ws://%s/ws/nodes/n1/backend-logs/dialed-model", echoAddr), nil) + Expect(err).ToNot(HaveOccurred()) + defer func() { _ = conn.Close() }() + + Expect(conn.SetReadDeadline(time.Now().Add(5 * time.Second))).To(Succeed()) + var initialMsg map[string]json.RawMessage + Expect(conn.ReadJSON(&initialMsg)).To(Succeed()) + var msgType string + Expect(json.Unmarshal(initialMsg["type"], &msgType)).To(Succeed()) + Expect(msgType).To(Equal("initial")) + Expect(dialedNodes()).To(Equal([]string{"n1", "n1", "n1"})) + }) + }) }) // startTestFileTransferServerWithLogs starts the real nodes.StartFileTransferServerWithListener diff --git a/tests/e2e/distributed/distributed_full_flow_test.go b/tests/e2e/distributed/distributed_full_flow_test.go index ad7f2669a..589fb2e41 100644 --- a/tests/e2e/distributed/distributed_full_flow_test.go +++ b/tests/e2e/distributed/distributed_full_flow_test.go @@ -13,6 +13,7 @@ import ( "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/grpc/base" pb "github.com/mudler/LocalAI/pkg/grpc/proto" @@ -256,12 +257,12 @@ var _ = Describe("Full Distributed Inference Flow", Label("Distributed"), func() // by registering after each node. In practice, we rely on the test registering // nodes before calling Route, so we subscribe to a catch-all pattern. infra.NC.Conn().Subscribe("nodes.*.backend.install", func(msg *nats.Msg) { - reply := messaging.BackendInstallReply{Success: true} + reply := workerctl.BackendInstallReply{Success: true} data, _ := json.Marshal(reply) msg.Respond(data) }) _, err := infra.NC.Conn().Subscribe("nodes.*.models.running", func(msg *nats.Msg) { - data, _ := json.Marshal(messaging.ModelsRunningReply{}) + data, _ := json.Marshal(workerctl.ModelsRunningReply{}) _ = msg.Respond(data) }) Expect(err).NotTo(HaveOccurred()) @@ -489,7 +490,7 @@ var _ = Describe("Full Distributed Inference Flow", Label("Distributed"), func() return "", err } return n.HTTPAddress, nil - }, "") + }, "", nodes.DirectWorkerNetDialer()) // Create SmartRouter with the HTTPFileStager router := newTestSmartRouter(registry, nodes.SmartRouterOptions{FileStager: stager}) @@ -558,7 +559,7 @@ var _ = Describe("Full Distributed Inference Flow", Label("Distributed"), func() return "", err } return n.HTTPAddress, nil - }, "") + }, "", nodes.DirectWorkerNetDialer()) // Create SmartRouter with FileStager router := newTestSmartRouter(registry, nodes.SmartRouterOptions{FileStager: stager}) @@ -616,7 +617,7 @@ var _ = Describe("Full Distributed Inference Flow", Label("Distributed"), func() return "", err } return n.HTTPAddress, nil - }, "") + }, "", nodes.DirectWorkerNetDialer()) // Test AllocRemoteTemp + FetchRemote directly (the output retrieval path) remoteTmpPath, err := stager.AllocRemoteTemp(ctx, node.ID) @@ -662,7 +663,7 @@ var _ = Describe("Full Distributed Inference Flow", Label("Distributed"), func() return "", err } return n.HTTPAddress, nil - }, "") + }, "", nodes.DirectWorkerNetDialer()) router := newTestSmartRouter(registry, nodes.SmartRouterOptions{FileStager: stager}) @@ -881,7 +882,7 @@ var _ = Describe("Full Distributed Inference Flow", Label("Distributed"), func() return "", err } return n.HTTPAddress, nil - }, "") + }, "", nodes.DirectWorkerNetDialer()) // Create model files on the "frontend" frontendModelsDir := GinkgoT().TempDir() @@ -965,7 +966,7 @@ var _ = Describe("Full Distributed Inference Flow", Label("Distributed"), func() return "", err } return n.HTTPAddress, nil - }, "") + }, "", nodes.DirectWorkerNetDialer()) // Create model files: .onnx and .onnx.json in a temp "models" dir frontendModelsDir := GinkgoT().TempDir() diff --git a/tests/e2e/distributed/file_staging_test.go b/tests/e2e/distributed/file_staging_test.go index 55bd5663c..e69f92087 100644 --- a/tests/e2e/distributed/file_staging_test.go +++ b/tests/e2e/distributed/file_staging_test.go @@ -62,7 +62,7 @@ var _ = Describe("File Staging", Label("Distributed"), func() { It("should create HTTPFileStager with httpAddrFor function", func() { stager := nodes.NewHTTPFileStager(func(nodeID string) (string, error) { return "", fmt.Errorf("no such node: %s", nodeID) - }, "") + }, "", nodes.DirectWorkerNetDialer()) Expect(stager).ToNot(BeNil()) // Should fail gracefully when node resolution fails diff --git a/tests/e2e/distributed/foundation_test.go b/tests/e2e/distributed/foundation_test.go index 244b5e6e0..944e66b6a 100644 --- a/tests/e2e/distributed/foundation_test.go +++ b/tests/e2e/distributed/foundation_test.go @@ -92,7 +92,7 @@ var _ = Describe("Phase 0: Foundation", Label("Distributed"), func() { Expect(client.IsConnected()).To(BeTrue()) received := make(chan []byte, 1) - sub, err := client.Subscribe("test.subject", func(data []byte) { + sub, err := client.Subscribe(messaging.SubjectJobProgress("e2e-pubsub"), func(data []byte) { received <- data }) Expect(err).ToNot(HaveOccurred()) @@ -101,7 +101,7 @@ var _ = Describe("Phase 0: Foundation", Label("Distributed"), func() { // Small delay to ensure subscription is active FlushNATS(client) - err = client.Publish("test.subject", map[string]string{"msg": "hello"}) + err = client.Publish(messaging.SubjectJobProgress("e2e-pubsub"), map[string]string{"msg": "hello"}) Expect(err).ToNot(HaveOccurred()) Eventually(received, "5s").Should(Receive()) @@ -114,13 +114,13 @@ var _ = Describe("Phase 0: Foundation", Label("Distributed"), func() { var worker1Count, worker2Count atomic.Int32 - sub1, err := client.QueueSubscribe("test.queue", "workers", func(data []byte) { + sub1, err := client.QueueSubscribe(messaging.SubjectJobProgress("e2e-queue"), "workers", func(data []byte) { worker1Count.Add(1) }) Expect(err).ToNot(HaveOccurred()) defer sub1.Unsubscribe() - sub2, err := client.QueueSubscribe("test.queue", "workers", func(data []byte) { + sub2, err := client.QueueSubscribe(messaging.SubjectJobProgress("e2e-queue"), "workers", func(data []byte) { worker2Count.Add(1) }) Expect(err).ToNot(HaveOccurred()) @@ -130,7 +130,7 @@ var _ = Describe("Phase 0: Foundation", Label("Distributed"), func() { // Publish multiple messages for i := range 10 { - err = client.Publish("test.queue", map[string]int{"n": i}) + err = client.Publish(messaging.SubjectJobProgress("e2e-queue"), map[string]int{"n": i}) Expect(err).ToNot(HaveOccurred()) } diff --git a/tests/e2e/distributed/job_dispatch_test.go b/tests/e2e/distributed/job_dispatch_test.go index 49052e593..a8274fe2c 100644 --- a/tests/e2e/distributed/job_dispatch_test.go +++ b/tests/e2e/distributed/job_dispatch_test.go @@ -7,6 +7,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/services/advisorylock" "github.com/mudler/LocalAI/core/services/jobs" + "github.com/mudler/LocalAI/core/services/messaging" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -38,9 +39,9 @@ var _ = Describe("Job Dispatch", Label("Distributed"), func() { Context("NATS job dispatch", func() { It("should enqueue job via NATS when dispatcher is set", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "dispatch-instance", 0) + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "dispatch-instance") var processed atomic.Int32 - dispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { processed.Add(1) store.UpdateJobStatus(job.ID, "completed", "done", "") return nil @@ -103,9 +104,9 @@ var _ = Describe("Job Dispatch", Label("Distributed"), func() { Context("NATS job cancellation", func() { It("should cancel running job via NATS cancel subject", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "cancel-instance", 0) + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "cancel-instance") jobStarted := make(chan struct{}) - dispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { close(jobStarted) <-ctx.Done() return ctx.Err() diff --git a/tests/e2e/distributed/job_distribution_test.go b/tests/e2e/distributed/job_distribution_test.go index fc6c3a031..d6ef1c879 100644 --- a/tests/e2e/distributed/job_distribution_test.go +++ b/tests/e2e/distributed/job_distribution_test.go @@ -3,6 +3,7 @@ package distributed_test import ( "context" "encoding/json" + "errors" "sync/atomic" "time" @@ -170,9 +171,9 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { Context("Job Distribution via NATS", func() { It("should enqueue job via NATS and worker picks it up", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "test-instance", 0) + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "test-instance") var processed atomic.Int32 - dispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { processed.Add(1) store.UpdateJobStatus(job.ID, "completed", "done", "") return nil @@ -201,9 +202,9 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { }) It("should cancel running job via NATS", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "test-instance", 0) + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "test-instance") jobStarted := make(chan struct{}) - dispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { close(jobStarted) // Simulate long work — wait for cancellation <-ctx.Done() @@ -239,8 +240,8 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { }) It("should report job progress via NATS", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "test-instance", 0) - dispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "test-instance") + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { dispatcher.PublishProgress(job.ID, "running", "step 1") time.Sleep(50 * time.Millisecond) dispatcher.PublishProgress(job.ID, "running", "step 2") @@ -317,7 +318,7 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { Context("Progress Streaming (NATS → SSE bridge)", func() { It("should bridge NATS progress events", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "test-instance", 0) + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "test-instance") dCtx, dCancel := context.WithCancel(infra.Ctx) defer dCancel() @@ -345,7 +346,7 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { }) It("should filter SSE events by job ID", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "test-instance", 0) + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "test-instance") dCtx, dCancel := context.WithCancel(infra.Ctx) defer dCancel() @@ -376,7 +377,7 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { Context("Enriched Job Payload (DB-free worker)", func() { It("should enrich JobEvent with full Job and Task data", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "enrichment-test", 0) + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "enrichment-test") dCtx, dCancel := context.WithCancel(infra.Ctx) defer dCancel() @@ -418,14 +419,12 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { }) It("should process job from enriched payload without DB access", func() { - // Create a worker-side dispatcher with NO store (simulating DB-free worker) - workerDispatcher := jobs.NewDispatcher(nil, infra.NC, nil, "worker-no-db", 0) - + // The worker has no store: everything it needs is in the payload. var receivedJob *jobs.JobRecord var receivedTask *jobs.TaskRecord processed := make(chan struct{}) - workerDispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { receivedJob = job receivedTask = task job.Result = "processed without DB" @@ -433,14 +432,7 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { return nil }) - dCtx, dCancel := context.WithCancel(infra.Ctx) - defer dCancel() - Expect(workerDispatcher.Start(dCtx)).To(Succeed()) - defer workerDispatcher.Stop() - - FlushNATS(infra.NC) - - // Publish an enriched event directly (simulating what the frontend does) + // Enqueue an enriched event directly (simulating what the frontend does) evt := jobs.JobEvent{ JobID: "test-job-123", TaskID: "test-task-456", @@ -459,7 +451,7 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { Prompt: "do something", }, } - Expect(infra.NC.Publish(messaging.SubjectJobsNew, evt)).To(Succeed()) + Expect(messaging.NewNATSWorkQueue(infra.NC).Enqueue(infra.Ctx, messaging.WorkTask, evt)).To(Succeed()) Eventually(processed, "10s").Should(BeClosed()) @@ -472,8 +464,8 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { }) It("should publish job result via NATS on completion", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "result-test", 0) - dispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "result-test") + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { job.Result = "job finished successfully" return nil }) @@ -507,8 +499,8 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { }) It("should stream traces via NATS progress events", func() { - dispatcher := jobs.NewDispatcher(store, infra.NC, db, "trace-test", 0) - dispatcher.SetWorkerFunc(func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { + dispatcher := jobs.NewDispatcher(store, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "trace-test") + startTaskWorker(infra.NC, func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error { dispatcher.PublishTrace(job.ID, "reasoning", "thinking about the problem") dispatcher.PublishTrace(job.ID, "tool_call", "calling search tool") return nil @@ -588,3 +580,57 @@ var _ = Describe("Phase 2: Jobs & Tasks", Label("Distributed"), func() { }) }) }) + +// startTaskWorker stands in for a task worker. Nothing in production consumes +// WorkTask, so these specs bring their own consumer on the WorkConsumer real +// workers use, and report the job lifecycle on events so the frontend +// dispatcher's result and progress subscriptions have something to persist. +func startTaskWorker(nc *messaging.Client, run func(ctx context.Context, job *jobs.JobRecord, task *jobs.TaskRecord) error) { + GinkgoHelper() + ctx, cancel := context.WithCancel(context.Background()) + sub, err := messaging.NewNATSWorkConsumer(nc).Consume(ctx, messaging.WorkTask, 0, func(ctx context.Context, payload []byte, events messaging.Publisher) error { + var evt jobs.JobEvent + if err := json.Unmarshal(payload, &evt); err != nil { + return err + } + if evt.Job == nil || evt.Task == nil { + jobs.PublishJobResult(events, evt.JobID, "failed", "", "job event carries no job or task") + return nil + } + + jobCtx, cancelJob := context.WithCancel(ctx) + defer cancelJob() + cancelSub, err := messaging.SubscribeJSON(nc, messaging.SubjectJobCancel(evt.JobID), func(jobs.CancelEvent) { + cancelJob() + }) + if err != nil { + return err + } + defer func() { _ = cancelSub.Unsubscribe() }() + // A spec cancels as soon as run signals it started; the cancel + // subscription has to be on the server by then. + if err := nc.Conn().Flush(); err != nil { + return err + } + + jobs.PublishJobProgress(events, evt.JobID, "running", "Job started") + runErr := run(jobCtx, evt.Job, evt.Task) + switch { + case errors.Is(jobCtx.Err(), context.Canceled): + jobs.PublishJobResult(events, evt.JobID, "cancelled", "", "") + case runErr != nil: + jobs.PublishJobResult(events, evt.JobID, "failed", "", runErr.Error()) + default: + jobs.PublishJobResult(events, evt.JobID, "completed", evt.Job.Result, "") + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + // Cancel first so a handler parked on its job context returns and the + // unsubscribe, which waits for in-flight handlers, does not hang. + DeferCleanup(func() { + cancel() + _ = sub.Unsubscribe() + }) + FlushNATS(nc) +} diff --git a/tests/e2e/distributed/managers_test.go b/tests/e2e/distributed/managers_test.go index b4f51ef95..dc4b3d712 100644 --- a/tests/e2e/distributed/managers_test.go +++ b/tests/e2e/distributed/managers_test.go @@ -12,6 +12,7 @@ import ( "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/system" @@ -139,21 +140,21 @@ var _ = Describe("Model and Backend Managers", Label("Distributed"), func() { // Subscribe to model.delete on both node subjects, track receipt var deleteCount atomic.Int32 sub1, err := infra.NC.SubscribeReply(messaging.SubjectNodeModelDelete(node1.ID), func(data []byte, reply func([]byte)) { - var req messaging.ModelDeleteRequest + var req workerctl.ModelDeleteRequest json.Unmarshal(data, &req) Expect(req.ModelName).To(Equal("big-model")) deleteCount.Add(1) - resp, _ := json.Marshal(messaging.ModelDeleteReply{Success: true}) + resp, _ := json.Marshal(workerctl.ModelDeleteReply{Success: true}) reply(resp) }) Expect(err).ToNot(HaveOccurred()) defer sub1.Unsubscribe() sub2, err := infra.NC.SubscribeReply(messaging.SubjectNodeModelDelete(node2.ID), func(data []byte, reply func([]byte)) { - var req messaging.ModelDeleteRequest + var req workerctl.ModelDeleteRequest json.Unmarshal(data, &req) deleteCount.Add(1) - resp, _ := json.Marshal(messaging.ModelDeleteReply{Success: true}) + resp, _ := json.Marshal(workerctl.ModelDeleteReply{Success: true}) reply(resp) }) Expect(err).ToNot(HaveOccurred()) @@ -205,21 +206,21 @@ var _ = Describe("Model and Backend Managers", Label("Distributed"), func() { // Subscribe to backend.delete on all 3 nodes var deleteCount atomic.Int32 sub1, err := infra.NC.SubscribeReply(messaging.SubjectNodeBackendDelete(node1.ID), func(data []byte, reply func([]byte)) { - var req messaging.BackendDeleteRequest + var req workerctl.BackendDeleteRequest json.Unmarshal(data, &req) Expect(req.Backend).To(Equal("my-backend")) deleteCount.Add(1) - resp, _ := json.Marshal(messaging.BackendDeleteReply{Success: true}) + resp, _ := json.Marshal(workerctl.BackendDeleteReply{Success: true}) reply(resp) }) Expect(err).ToNot(HaveOccurred()) defer sub1.Unsubscribe() sub2, err := infra.NC.SubscribeReply(messaging.SubjectNodeBackendDelete(node2.ID), func(data []byte, reply func([]byte)) { - var req messaging.BackendDeleteRequest + var req workerctl.BackendDeleteRequest json.Unmarshal(data, &req) deleteCount.Add(1) - resp, _ := json.Marshal(messaging.BackendDeleteReply{Success: true}) + resp, _ := json.Marshal(workerctl.BackendDeleteReply{Success: true}) reply(resp) }) Expect(err).ToNot(HaveOccurred()) @@ -228,7 +229,7 @@ var _ = Describe("Model and Backend Managers", Label("Distributed"), func() { var unhealthyReceived atomic.Int32 sub3, err := infra.NC.SubscribeReply(messaging.SubjectNodeBackendDelete(node3.ID), func(data []byte, reply func([]byte)) { unhealthyReceived.Add(1) - resp, _ := json.Marshal(messaging.BackendDeleteReply{Success: true}) + resp, _ := json.Marshal(workerctl.BackendDeleteReply{Success: true}) reply(resp) }) Expect(err).ToNot(HaveOccurred()) @@ -276,11 +277,11 @@ var _ = Describe("Model and Backend Managers", Label("Distributed"), func() { var deleteCount atomic.Int32 sub1, err := infra.NC.SubscribeReply(messaging.SubjectNodeBackendDelete(node1.ID), func(data []byte, reply func([]byte)) { - var req messaging.BackendDeleteRequest + var req workerctl.BackendDeleteRequest json.Unmarshal(data, &req) Expect(req.Backend).To(Equal("remote-only-backend")) deleteCount.Add(1) - resp, _ := json.Marshal(messaging.BackendDeleteReply{Success: true}) + resp, _ := json.Marshal(workerctl.BackendDeleteReply{Success: true}) reply(resp) }) Expect(err).ToNot(HaveOccurred()) diff --git a/tests/e2e/distributed/mcp_nats_test.go b/tests/e2e/distributed/mcp_nats_test.go index e0e868e4d..196a1a5a0 100644 --- a/tests/e2e/distributed/mcp_nats_test.go +++ b/tests/e2e/distributed/mcp_nats_test.go @@ -1,7 +1,9 @@ package distributed_test import ( + "context" "encoding/json" + "strings" "sync/atomic" "time" @@ -9,6 +11,7 @@ import ( mcpTools "github.com/mudler/LocalAI/core/http/endpoints/mcp" mcpRemote "github.com/mudler/LocalAI/core/services/mcp" "github.com/mudler/LocalAI/core/services/messaging" + "github.com/mudler/LocalAI/core/services/nodes" "github.com/mudler/LocalAI/pkg/functions" . "github.com/onsi/ginkgo/v2" @@ -47,7 +50,7 @@ var _ = Describe("MCP NATS Routing", Label("Distributed"), func() { // Frontend side: pass NATS client and call remote result, err := mcpTools.ExecuteMCPToolCallRemote( infra.Ctx, - infra.NC, + nodes.NewNATSAgentControl(infra.NC), "test-model", config.MCPGenericConfig[config.MCPRemoteServers]{}, config.MCPGenericConfig[config.MCPSTDIOServers]{}, @@ -72,7 +75,7 @@ var _ = Describe("MCP NATS Routing", Label("Distributed"), func() { _, err = mcpTools.ExecuteMCPToolCallRemote( infra.Ctx, - infra.NC, + nodes.NewNATSAgentControl(infra.NC), "test-model", config.MCPGenericConfig[config.MCPRemoteServers]{}, config.MCPGenericConfig[config.MCPSTDIOServers]{}, @@ -111,7 +114,7 @@ var _ = Describe("MCP NATS Routing", Label("Distributed"), func() { result, err := mcpTools.DiscoverMCPToolsRemote( infra.Ctx, - infra.NC, + nodes.NewNATSAgentControl(infra.NC), "discovery-model", config.MCPGenericConfig[config.MCPRemoteServers]{}, config.MCPGenericConfig[config.MCPSTDIOServers]{}, @@ -125,10 +128,66 @@ var _ = Describe("MCP NATS Routing", Label("Distributed"), func() { }) }) + Context("Agent RPC server", func() { + It("round trips tool and discovery requests through the agent RPC server", func() { + toolReqs := make(chan mcpRemote.MCPToolRequest, 2) + discoveryReqs := make(chan mcpRemote.MCPDiscoveryRequest, 1) + + srv := nodes.NewNATSAgentRPCServer(infra.NC, "e2e-agent-node") + Expect(srv.ServeMCPTool(func(_ context.Context, req mcpRemote.MCPToolRequest) mcpRemote.MCPToolResponse { + toolReqs <- req + return mcpRemote.MCPToolResponse{Result: "ran " + req.ToolName} + })).To(Succeed()) + Expect(srv.ServeMCPDiscovery(func(_ context.Context, req mcpRemote.MCPDiscoveryRequest) mcpRemote.MCPDiscoveryResponse { + discoveryReqs <- req + return mcpRemote.MCPDiscoveryResponse{ + Servers: []mcpRemote.MCPServerInfo{{Name: "weather-server", Type: "remote", Tools: []string{"get_weather"}}}, + } + })).To(Succeed()) + FlushNATS(infra.NC) + + control := nodes.NewNATSAgentControl(infra.NC) + + toolResp, err := control.ExecuteMCPTool(infra.Ctx, mcpRemote.MCPToolRequest{ + ModelName: "rpc-model", + ToolName: "get_weather", + Arguments: map[string]any{"city": "Rome"}, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(toolResp.Result).To(Equal("ran get_weather")) + Expect(toolResp.Error).To(BeEmpty()) + var gotTool mcpRemote.MCPToolRequest + Eventually(toolReqs).Should(Receive(&gotTool)) + Expect(gotTool.ModelName).To(Equal("rpc-model")) + Expect(gotTool.ToolName).To(Equal("get_weather")) + Expect(gotTool.Arguments).To(HaveKeyWithValue("city", "Rome")) + + discoveryResp, err := control.DiscoverMCPTools(infra.Ctx, mcpRemote.MCPDiscoveryRequest{ModelName: "rpc-model"}) + Expect(err).ToNot(HaveOccurred()) + Expect(discoveryResp.Servers).To(HaveLen(1)) + Expect(discoveryResp.Servers[0].Name).To(Equal("weather-server")) + Expect(discoveryResp.Servers[0].Tools).To(ConsistOf("get_weather")) + var gotDiscovery mcpRemote.MCPDiscoveryRequest + Eventually(discoveryReqs).Should(Receive(&gotDiscovery)) + Expect(gotDiscovery.ModelName).To(Equal("rpc-model")) + + // AgentControl only sends valid JSON, so the undecodable body goes + // on the wire directly: the server must still answer, or the + // requester would wait out its whole budget. + raw, err := infra.NC.Request(messaging.SubjectMCPToolExecute, []byte("{not json"), 5*time.Second) + Expect(err).ToNot(HaveOccurred()) + var refused mcpRemote.MCPToolResponse + Expect(json.Unmarshal(raw, &refused)).To(Succeed()) + Expect(strings.HasPrefix(refused.Error, "unmarshal error: ")).To(BeTrue(), "got %q", refused.Error) + Expect(refused.Result).To(BeEmpty()) + Consistently(toolReqs, 200*time.Millisecond).ShouldNot(Receive()) + }) + }) + Context("QueueSubscribeReply", func() { It("should support queue subscribe with request-reply round-trip", func() { // Subscribe with queue group - sub, err := infra.NC.QueueSubscribeReply("test.echo", "echo-workers", func(data []byte, reply func([]byte)) { + sub, err := infra.NC.QueueSubscribeReply(messaging.SubjectNodeBackendList("e2e-echo"), "echo-workers", func(data []byte, reply func([]byte)) { // Echo back the request data with a prefix reply(append([]byte("echo:"), data...)) }) @@ -138,7 +197,7 @@ var _ = Describe("MCP NATS Routing", Label("Distributed"), func() { FlushNATS(infra.NC) // Send request and wait for reply - replyData, err := infra.NC.Request("test.echo", []byte("hello"), 5*time.Second) + replyData, err := infra.NC.Request(messaging.SubjectNodeBackendList("e2e-echo"), []byte("hello"), 5*time.Second) Expect(err).ToNot(HaveOccurred()) Expect(string(replyData)).To(Equal("echo:hello")) }) @@ -146,13 +205,13 @@ var _ = Describe("MCP NATS Routing", Label("Distributed"), func() { It("should load-balance requests across queue subscribers", func() { var worker1Count, worker2Count atomic.Int32 - sub1, _ := infra.NC.QueueSubscribeReply("test.lb", "lb-workers", func(data []byte, reply func([]byte)) { + sub1, _ := infra.NC.QueueSubscribeReply(messaging.SubjectNodeBackendList("e2e-lb"), "lb-workers", func(data []byte, reply func([]byte)) { worker1Count.Add(1) reply([]byte("w1")) }) defer sub1.Unsubscribe() - sub2, _ := infra.NC.QueueSubscribeReply("test.lb", "lb-workers", func(data []byte, reply func([]byte)) { + sub2, _ := infra.NC.QueueSubscribeReply(messaging.SubjectNodeBackendList("e2e-lb"), "lb-workers", func(data []byte, reply func([]byte)) { worker2Count.Add(1) reply([]byte("w2")) }) @@ -162,7 +221,7 @@ var _ = Describe("MCP NATS Routing", Label("Distributed"), func() { // Send multiple requests for range 10 { - _, err := infra.NC.Request("test.lb", []byte("req"), 5*time.Second) + _, err := infra.NC.Request(messaging.SubjectNodeBackendList("e2e-lb"), []byte("req"), 5*time.Second) Expect(err).ToNot(HaveOccurred()) } diff --git a/tests/e2e/distributed/model_config_revision_test.go b/tests/e2e/distributed/model_config_revision_test.go index 548a35eb4..1b2c85c97 100644 --- a/tests/e2e/distributed/model_config_revision_test.go +++ b/tests/e2e/distributed/model_config_revision_test.go @@ -6,8 +6,8 @@ import ( "sync" "github.com/mudler/LocalAI/core/config" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/workerctl" pb "github.com/mudler/LocalAI/pkg/grpc/proto" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -22,14 +22,14 @@ type revisionCleanupStopper struct { stopped []nodes.NodeModel } -func (s *revisionCleanupStopper) StopModelReplica(_ context.Context, nodeID string, replica nodes.NodeModel, _ bool) (messaging.ModelStopReply, error) { +func (s *revisionCleanupStopper) StopModelReplica(_ context.Context, nodeID string, replica nodes.NodeModel, _ bool) (workerctl.ModelStopReply, error) { s.mu.Lock() defer s.mu.Unlock() s.stopped = append(s.stopped, replica) if nodeID == s.unreachable { - return messaging.ModelStopReply{}, errors.New("worker unreachable") + return workerctl.ModelStopReply{}, errors.New("worker unreachable") } - return messaging.ModelStopReply{ + return workerctl.ModelStopReply{ Matched: true, Terminated: true, ProcessKey: replica.ModelName, diff --git a/tests/e2e/distributed/nats_jwt_test.go b/tests/e2e/distributed/nats_jwt_test.go index bf947e472..27885ae1c 100644 --- a/tests/e2e/distributed/nats_jwt_test.go +++ b/tests/e2e/distributed/nats_jwt_test.go @@ -26,8 +26,9 @@ var _ = Describe("NATS JWT Auth", Label("Distributed", "NatsJWT"), func() { }) It("allows backend subscribe on the node prefix", func() { - wild := nodeSubjectPrefix(infra.NodeID) + ".>" - sub, err := infra.NC.Subscribe(wild, func(_ []byte) {}) + // The client refuses a `>` filter, so probe the prefix grant with a + // concrete subject under it rather than the wildcard itself. + sub, err := infra.NC.Subscribe(messaging.SubjectNodeBackendInstall(infra.NodeID), func(_ []byte) {}) Expect(err).ToNot(HaveOccurred()) defer func() { _ = sub.Unsubscribe() }() Expect(infra.NC.Conn().FlushTimeout(2 * time.Second)).To(Succeed()) diff --git a/tests/e2e/distributed/node_lifecycle_test.go b/tests/e2e/distributed/node_lifecycle_test.go index 04b7342e7..2ca7707c1 100644 --- a/tests/e2e/distributed/node_lifecycle_test.go +++ b/tests/e2e/distributed/node_lifecycle_test.go @@ -8,6 +8,7 @@ import ( "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/workerctl" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -46,11 +47,11 @@ var _ = Describe("Node Backend Lifecycle (NATS-driven)", Label("Distributed"), f // Simulate worker subscribing to backend.install and replying success infra.NC.SubscribeReply(messaging.SubjectNodeBackendInstall(node.ID), func(data []byte, reply func([]byte)) { - var req messaging.BackendInstallRequest + var req workerctl.BackendInstallRequest json.Unmarshal(data, &req) Expect(req.Backend).To(Equal("llama-cpp")) - resp := messaging.BackendInstallReply{Success: true} + resp := workerctl.BackendInstallReply{Success: true} respData, _ := json.Marshal(resp) reply(respData) }) @@ -71,7 +72,7 @@ var _ = Describe("Node Backend Lifecycle (NATS-driven)", Label("Distributed"), f // Simulate worker replying with error infra.NC.SubscribeReply(messaging.SubjectNodeBackendInstall(node.ID), func(data []byte, reply func([]byte)) { - resp := messaging.BackendInstallReply{Success: false, Error: "backend not found"} + resp := workerctl.BackendInstallReply{Success: false, Error: "backend not found"} respData, _ := json.Marshal(resp) reply(respData) }) diff --git a/tests/e2e/distributed/prefix_cache_routing_test.go b/tests/e2e/distributed/prefix_cache_routing_test.go index 9b1e3c117..852fdb5e1 100644 --- a/tests/e2e/distributed/prefix_cache_routing_test.go +++ b/tests/e2e/distributed/prefix_cache_routing_test.go @@ -47,7 +47,7 @@ type prefixStubClientFactory struct { client *prefixStubBackend } -func (f *prefixStubClientFactory) NewClient(_ string, _ bool) grpcPkg.Backend { +func (f *prefixStubClientFactory) NewClient(_, _ string, _ bool) grpcPkg.Backend { return f.client } diff --git a/tests/e2e/distributed/router_tracking_test.go b/tests/e2e/distributed/router_tracking_test.go index 75895a372..50a3659a0 100644 --- a/tests/e2e/distributed/router_tracking_test.go +++ b/tests/e2e/distributed/router_tracking_test.go @@ -5,8 +5,8 @@ import ( "encoding/json" "time" - "github.com/mudler/LocalAI/core/services/messaging" "github.com/mudler/LocalAI/core/services/nodes" + "github.com/mudler/LocalAI/core/services/workerctl" "github.com/mudler/LocalAI/pkg/grpc/base" pb "github.com/mudler/LocalAI/pkg/grpc/proto" @@ -62,12 +62,12 @@ var _ = Describe("SmartRouter trackingKey", Label("Distributed"), func() { // Mock backend.install handler — always replies success infra.NC.Conn().Subscribe("nodes.*.backend.install", func(msg *nats.Msg) { - reply := messaging.BackendInstallReply{Success: true} + reply := workerctl.BackendInstallReply{Success: true} data, _ := json.Marshal(reply) msg.Respond(data) }) _, err = infra.NC.Conn().Subscribe("nodes.*.models.running", func(msg *nats.Msg) { - data, _ := json.Marshal(messaging.ModelsRunningReply{}) + data, _ := json.Marshal(workerctl.ModelsRunningReply{}) _ = msg.Respond(data) }) Expect(err).NotTo(HaveOccurred()) diff --git a/tests/e2e/distributed/sse_routes_test.go b/tests/e2e/distributed/sse_routes_test.go index 4cc334814..2be9b3d68 100644 --- a/tests/e2e/distributed/sse_routes_test.go +++ b/tests/e2e/distributed/sse_routes_test.go @@ -6,6 +6,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/services/agents" "github.com/mudler/LocalAI/core/services/jobs" + "github.com/mudler/LocalAI/core/services/messaging" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" @@ -36,7 +37,7 @@ var _ = Describe("SSE Routes", Label("Distributed"), func() { jobStore, err := jobs.NewJobStore(db) Expect(err).ToNot(HaveOccurred()) - dispatcher := jobs.NewDispatcher(jobStore, infra.NC, db, "sse-instance", 0) + dispatcher := jobs.NewDispatcher(jobStore, messaging.NewNATSWorkQueue(infra.NC), infra.NC, db, "sse-instance") dCtx, dCancel := context.WithCancel(infra.Ctx) defer dCancel()