From e4fa051ee32e916c62ee7c5b26c88ab8decc568a Mon Sep 17 00:00:00 2001 From: mudler-agent Date: Fri, 2 Oct 2026 23:58:41 +0200 Subject: [PATCH] refactor(distributed): put the NATS-only paths behind interfaces (#12395) * feat(messaging): add shared subject rules Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * test(messaging): cover BroadcastRoots, ControlRoots and SubjectRoot Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * feat(messaging): add Broadcaster and enforce subject rules in every carrier Broadcaster is the fan-out half of MessagingClient. The NATS client and the in-memory FakeBus now refuse a subject outside the served roots and any wildcard other than a whole single token, and FakeBus shares MatchSubject instead of its own copy. FakeBus Unsubscribe now removes its own subscription instead of the first one with the same subject. A shared conformance suite in messagingtest runs against both carriers. The distributed e2e specs that used invented test.* subjects, and the one that subscribed with a > filter, now use subjects from subjects.go. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor: depend on Broadcaster where only publish and subscribe are used Narrowed to messaging.Broadcaster: nodes/staging_progress.go, nodes/install_progress_publisher.go, galleryop/operation.go, galleryop/service.go, agentpool/user_services.go, agentpool/agent_jobs.go, openresponses/store.go, openresponses/sync.go, syncstate/syncstate.go, finetune/service.go, quantization/service.go and failover/distsync/distsync.go. SubscribeJSON now takes a Broadcaster because it only calls Subscribe, which lets the narrowed consumers use it. Stayed wide: worker/supervisor.go, because its client field also serves the SubscribeReply handlers in worker/lifecycle.go. The request/reply, queue and wiring files (nodes/unloader.go, nodes/file_stager_s3.go, jobs/dispatcher.go, agents/dispatcher.go, agents/events.go, worker/file_staging.go, cli/agent_worker.go, http/app.go) are unchanged by design. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(nodes): name the no-route condition and confine the carrier error Consumers matched nats.ErrNoResponders, which names an absence, to demote a node. They now match ErrNoRoute, the control path maps the carrier's failure onto it, and timeouts and worker refusals are pinned as not being no-route. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * docs(nodes): state which FileStager implementations return ErrNoRoute Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(nodes): build backend clients through one node-aware seam Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * docs: describe the distributed transport seams Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * docs: correct comments that overclaim after the seams refactor Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * fix(agent-worker): refuse an unserved LOCALAI_AGENT_SUBJECT at startup The messaging client now refuses a subject whose root no carrier serves. An agent worker started with a custom LOCALAI_AGENT_SUBJECT such as tenant-a.agent.execute used to start and then wait on a subject the frontend never publishes to. After the subject rules landed it exited at subscribe time with an error that did not name the setting. Behaviour change: the worker now checks LOCALAI_AGENT_SUBJECT before it registers or connects, and exits with an error that names the variable and says to use a served subject under the agent root, for example agent.execute. The served roots are not widened: a custom root was never delivered by the frontend, and a wider set would reopen the drift the subject rules exist to close. The flag help and the agent worker docs state the constraint. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * test(nodes): pin the reactions to ErrNoRoute Three callers react to ErrNoRoute and had no spec: the reconciler's upgrade drain falls back to the legacy forced install, the reconciler marks the node unhealthy when a pending op has no route, and the backend-op fan-out marks the node unhealthy. Each spec drives the real caller with a scripted no-responders reply and reads the result from the registry or the recorded requests. A fourth spec pins the other side: a pending op that times out leaves the node healthy and only counts the attempt, so mapping timeouts onto ErrNoRoute would fail here. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * test(messaging): pin client subject checks and fail the carrier suite in CI Add specs that call Publish, Request, Subscribe, QueueSubscribe, SubscribeReply and QueueSubscribeReply on a client with no connection. Each call must return ErrUnservedSubject for bogus.thing and ErrUnsupportedWildcard for jobs.>. This proves that the subject check runs before the connection is used, and needs no server. The NATS conformance suite is the only check that runs the subject rules against a real carrier. Before this change it skipped without output when Docker was missing. Now it fails when CI is set, so a Linux runner without Docker cannot hide it. It still skips on local runs and on macOS CI, which has no Docker. Add SubjectNodeBackendInstallProgress to the list of constructors that must build served subjects, and ask contributors to extend the list. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * docs: state what ErrNoRoute may change, and group the distributed guides The seams note said MarkUnhealthy was the only state change allowed on ErrNoRoute. A pending backend op still records the failed attempt, counts toward the reconciler's retry limit and is dead-lettered after the maximum attempts. The note now says that MarkUnhealthy is the only change to the node's own state, and that the per-op accounting is not a verdict about the node. The note also documents that the NATS conformance run fails under CI when Docker is missing. The distributed-seams row moves next to the distributed-state row in the topics table. The liveness ping spec header now says no route is a reason to skip the worker, not proof that the worker is gone. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(nodes): give the backend client factory the node id Mechanical: the method gains a nodeID parameter and the eight test fakes are updated. No behaviour change. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(nodes): drop the optional node-aware factory The node id is now in the main method, so the optional interface and its helper had no behaviour of their own. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(nodes): dial backend probes through the client factory Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(nodes): dial workers' file servers through a per-node dialer Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(http): proxy backend logs through the per-node worker dialer The admin backend-logs proxy (list, lines and the WebSocket stream) now reaches a worker through the same per-node dialer as the HTTP file stager, so every frontend-to-worker dial goes through one seam. The shared direct dialer keeps alive for 15s where the proxy used 30s. Harmless for requests bounded at 15s. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * fix(http): keep the backend-logs proxy independent of the admin connection The proxy request had no context before the dialer change and is bounded only by its 15s timeout. Keep it that way. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor: move the worker control payloads to workerctl Mechanical move of the request and reply structs, the install progress event and the file payloads out of messaging. The verbs no longer belong to one carrier. No alias is left behind. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(worker): serve the lifecycle verbs through a controlServer The worker registers one handler per verb and a NATS server maps each verb to its subject. Registration errors now name the verb. node.stop is served with SubscribeReply, which is identical on the wire because the handler never replies. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(worker): report install progress through the control sink Install and upgrade now emit download progress through the sink the control server hands them. The debounce and the terminal flush stay in the handler path, built over that sink by the new nodes.NewDebouncedInstallProgressSink, which replaces NewDebouncedInstallProgressPublisher. The subject and payload on the wire are unchanged. The supervisor no longer holds the bus, and installFn and upgradeFn let specs drive both verbs without a gallery. The malformed-request log lines are restored for install, upgrade, backend.delete, model.unload, model.stop and model.delete, with the reply bytes unchanged. The signal adapter is renamed noReply, which also lets worker.go import os/signal without an alias again. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(worker): serve the file-staging verbs through a controlServer An empty list-dir answer is now {} rather than {"files":null}, because the typed reply omits an empty Files slice. The frontend decodes both to a nil slice in nodes/file_stager_s3.go. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * feat(messaging): add WorkQueue and the NATS producer This is the producer side of the competing-consumer seam. The work kinds map one to one to today's subjects and queue groups: task to jobs.new and mcp-ci to jobs.mcp-ci.new (both in group workers), agent-run to agent.execute (group agent-workers). Enqueue publishes the payload as Publish does today, with one JSON marshal. FakeBus now records queue groups and keeps reply handlers so later specs can pin and drive them. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * feat(messaging): add the NATS WorkConsumer An in-flight limit of one runs the handler inline on the delivery goroutine, as the MCP CI consumer does today. Any other limit spawns per delivery, as the agent consumer does. Queue groups are unchanged. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor: publish queued work through WorkQueue The job dispatcher, the agent pool and the agent scheduler enqueue through messaging.WorkQueue; the NATS implementation publishes to the same subjects as before. DistributedServices builds the queue next to the NATS client and hands it to the dispatcher and the agent pool, whose distributed mode switch now reads a non-nil WorkQueue. The unused AgentPoolService.SetNATSClient is removed. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor: consume queued work through WorkConsumer The agent dispatcher and the MCP CI consumer register through messaging.WorkConsumer. The NATS implementation keeps the inline one-at-a-time model for MCP CI and the per-delivery model for agent runs. handleMCPCIJob reports on the events publisher the carrier hands it instead of a captured client. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor: delete the consumers nothing in production reached jobs.new has a producer and no production consumer, and the agent dispatcher's Dispatch was only called from tests. Publishing jobs.new is unchanged. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(mcp): send MCP requests to agent workers through AgentControl Timeouts still honour only the deadline, not cancellation, exactly as today. The NATS no-responders error maps to ErrNoRoute and a timeout does not. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(agent-worker): serve MCP requests and backend.stop through agentRPCServer The agent worker's MCP tool and discovery reply subscriptions and its backend stop listener move behind an unexported agentRPCServer interface, served on NATS by nodes.NATSAgentRPCServer. The handlers become typed mcp.ToolHandler and mcp.DiscoveryHandler values that answer every failure with a reply carrying Error. Queue group (agent-workers), inline execution on the delivery goroutine, the background handler context, the unmarshal error reply texts and the reply-less backend stop subscription are unchanged. The backend stop handler takes the decoded backend name, so it can still close that backend's MCP sessions. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor(messaging): remove helpers that only tests used BroadcastRoots, ControlRoots and SubjectRoot had no production caller. The roots spec now asserts every served root through ValidateSubject instead. MatchSubject moves back into the test support package, the only place that used it, with its table. NATSAgentRPCServer drops the subscription list it stored and never read, and NewNATSAgentRPCServer gets a doc comment. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * test(mcp): round trip the agent RPC server over a real NATS server One spec sends a tool request and a discovery request through NATSAgentControl to NATSAgentRPCServer and checks that the handlers see the decoded requests and the replies come back. It also puts an undecodable body on the tool subject and checks the server answers with an unmarshal error instead of leaving the requester to time out. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * docs: describe the distributed transport seams The developer note now lists the final seams: fan-out, queues, both halves of the control verbs and of agent RPC, and the dial. It records the open items a second carrier has to handle. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * test: pin the in-flight limit each queue consumer asks for The work queue specs pin what Consume does for a given limit, but nothing pinned which limit each production consumer passes. Changing the agent worker's MCP CI limit from 1 to 0 would have let MCP CI jobs run concurrently on each worker with every test green. Move the MCP CI Consume call into startMCPCIConsumer with the same wiring and pin that it asks for (WorkMCPCI, 1). Pin that NATSDispatcher.Start asks for (WorkAgentRun, maxConcurrent) for several limits. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * refactor: remove helpers the branch left without a caller SubjectJobCancelWildcard lost its last subscriber when the frontend stopped listening on jobs.*.cancel; the NATS permissions and conformance suite spell the subject out, so nothing reads the constant. decodeBackendStopRequest returned a stopAll flag that production dropped and only a test read. decodeBackendStop is now the single decoder with the same semantics: an empty body is stop-all, an empty Backend is stop-all, malformed JSON is an error. stopBackends still derives stop-all from Backend, so no reply changes. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * fix(messaging): keep an explicitly empty agent queue a plain subscription Before the work queue seam the agent worker passed LOCALAI_AGENT_QUEUE straight to QueueSubscribe, so an explicitly empty value made a plain subscription and every agent worker ran every agent run. WithAgentRunRoute replaced an empty queue with agent-workers, which silently changed that. Keep the queue as given once the option is applied. An empty subject still falls back to agent.execute, since it never had a meaning of its own. The flag default stays agent-workers, so only an explicitly empty value reaches this. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto * docs: correct comments and record the PR B notes Fix the recordingFactory comment (it also records the parallel flag), document that a negative maxInFlight is unbounded and that Unsubscribe from a handler deadlocks, and say a permanently undecodable payload returns nil. Record controlHandler's undecodable return as a kept exception, and add the second carrier notes to the developer note: the reconciler has no ClientFactory option, the logs proxy honours HTTP_PROXY, verbs one carrier serves need an opt-out, terminal replies come from the result event, and agent runs publish through the NATS-bound EventBridge, which is not an additive change. Assisted-by: Claude:claude-sonnet-5-5 Signed-off-by: Ettore Di Giacinto --------- Signed-off-by: Ettore Di Giacinto Co-authored-by: Ettore Di Giacinto --- .agents/distributed-seams.md | 169 ++++++++ AGENTS.md | 1 + core/application/agent_pool_options_test.go | 39 ++ core/application/application.go | 25 +- core/application/distributed.go | 18 +- core/application/startup.go | 3 +- core/cli/agent_worker.go | 200 +++++---- core/cli/agent_worker_mcp_rpc_test.go | 47 ++ core/cli/agent_worker_mcpci_test.go | 139 ++++++ core/cli/agent_worker_subject_test.go | 31 ++ core/config/distributed_config.go | 1 - core/http/app.go | 4 +- core/http/endpoints/anthropic/messages.go | 4 +- core/http/endpoints/localai/mcp.go | 4 +- core/http/endpoints/localai/mcp_tools.go | 12 +- core/http/endpoints/localai/nodes.go | 29 +- .../localai/nodes_backends_list_test.go | 22 +- core/http/endpoints/mcp/executor.go | 40 +- .../mcp/executor_agent_control_test.go | 76 ++++ core/http/endpoints/mcp/tools.go | 54 ++- core/http/endpoints/openai/chat.go | 4 +- .../http/endpoints/openresponses/responses.go | 4 +- core/http/endpoints/openresponses/store.go | 2 +- core/http/endpoints/openresponses/sync.go | 4 +- core/http/routes/anthropic.go | 6 +- core/http/routes/localai.go | 8 +- core/http/routes/nodes.go | 8 +- core/http/routes/openai.go | 8 +- core/http/routes/openresponses.go | 8 +- core/services/agentpool/agent_jobs.go | 4 +- core/services/agentpool/agent_pool.go | 26 +- core/services/agentpool/dispatch_chat_test.go | 65 +++ core/services/agentpool/user_services.go | 4 +- core/services/agents/dispatcher.go | 121 ++---- core/services/agents/dispatcher_test.go | 159 +++++++ core/services/agents/scheduler.go | 23 +- core/services/agents/scheduler_test.go | 76 ++-- core/services/failover/distsync/distsync.go | 2 +- core/services/finetune/service.go | 2 +- core/services/galleryop/operation.go | 4 +- core/services/galleryop/service.go | 8 +- core/services/jobs/dispatcher.go | 217 +--------- core/services/jobs/dispatcher_test.go | 181 +++++--- core/services/mcp/remote.go | 11 + .../messaging/backend_install_progress.go | 28 -- .../backend_install_progress_test.go | 37 -- core/services/messaging/client.go | 24 +- .../messaging/client_conformance_test.go | 81 ++++ .../messaging/client_validation_test.go | 56 +++ core/services/messaging/export_test.go | 5 + core/services/messaging/interfaces.go | 16 +- core/services/messaging/interfaces_test.go | 14 + .../messaging/messagingtest/broadcaster.go | 154 +++++++ core/services/messaging/subject_rules.go | 64 +++ core/services/messaging/subject_rules_test.go | 100 +++++ core/services/messaging/subjects.go | 234 +--------- .../messaging/subjects_upgrade_test.go | 12 - core/services/messaging/workqueue.go | 45 ++ core/services/messaging/workqueue_nats.go | 196 +++++++++ .../services/messaging/workqueue_nats_test.go | 346 +++++++++++++++ core/services/nodes/agent_rpc_nats.go | 132 ++++++ core/services/nodes/agent_rpc_nats_test.go | 388 +++++++++++++++++ core/services/nodes/client_factory_test.go | 81 ++++ core/services/nodes/control_errors_test.go | 63 +++ core/services/nodes/control_nats.go | 59 +++ core/services/nodes/disk_headroom_test.go | 4 +- core/services/nodes/file_stager.go | 4 + core/services/nodes/file_stager_dial_test.go | 101 +++++ core/services/nodes/file_stager_http.go | 89 ++-- .../nodes/file_stager_release_test.go | 13 +- core/services/nodes/file_stager_s3.go | 65 +-- .../nodes/file_stager_verify_deadline_test.go | 2 +- .../file_staging_sound_detection_test.go | 2 +- .../nodes/file_transfer_finalize_test.go | 2 +- .../nodes/file_transfer_server_test.go | 20 +- core/services/nodes/health.go | 2 +- core/services/nodes/health_mock_test.go | 2 +- .../nodes/install_progress_publisher.go | 93 ++-- .../nodes/install_progress_publisher_test.go | 73 ++-- core/services/nodes/interfaces.go | 27 +- core/services/nodes/managers_distributed.go | 31 +- .../nodes/managers_distributed_test.go | 133 ++---- core/services/nodes/model_cleanup_test.go | 22 +- core/services/nodes/noroute_reactions_test.go | 161 +++++++ .../services/nodes/pending_op_cleanup_test.go | 2 +- core/services/nodes/reconciler.go | 36 +- core/services/nodes/reconciler_prober_test.go | 40 ++ core/services/nodes/reconciler_test.go | 32 +- .../nodes/reconciler_worker_processes_test.go | 22 +- core/services/nodes/registry.go | 6 +- .../nodes/revision_eligibility_test.go | 17 +- core/services/nodes/router.go | 2 +- core/services/nodes/router_liveness.go | 5 +- .../nodes/router_liveness_route_test.go | 44 ++ .../services/nodes/router_load_budget_test.go | 6 +- core/services/nodes/router_load_job_test.go | 6 +- .../nodes/router_load_timeout_test.go | 6 +- core/services/nodes/router_reap_load_test.go | 6 +- .../nodes/router_revision_lifecycle_test.go | 8 +- .../nodes/router_slot_uncertainty_test.go | 6 +- .../nodes/router_staging_context_test.go | 4 +- .../nodes/router_staging_deadline_test.go | 4 +- core/services/nodes/router_test.go | 37 +- core/services/nodes/staging_progress.go | 2 +- core/services/nodes/unloader.go | 97 ++--- core/services/nodes/unloader_ping_test.go | 10 +- .../nodes/unloader_stale_rows_test.go | 5 +- core/services/nodes/unloader_test.go | 39 +- core/services/nodes/unloader_upgrade_test.go | 11 +- core/services/quantization/service.go | 2 +- core/services/syncstate/syncstate.go | 4 +- core/services/testutil/export_test.go | 4 + core/services/testutil/fakebus.go | 122 ++++-- .../testutil/fakebus_conformance_test.go | 15 + core/services/testutil/fakebus_queue_test.go | 39 ++ core/services/testutil/subject_match_test.go | 21 + core/services/testutil/testutil_suite_test.go | 13 + core/services/worker/control_nats.go | 99 +++++ core/services/worker/control_nats_test.go | 374 ++++++++++++++++ core/services/worker/control_server.go | 110 +++++ core/services/worker/file_staging.go | 322 +++++++------- .../worker/file_staging_release_test.go | 8 +- .../worker/file_staging_verbs_test.go | 157 +++++++ core/services/worker/install.go | 38 +- core/services/worker/lifecycle.go | 409 +++++++++--------- core/services/worker/model_stop_test.go | 27 +- core/services/worker/models_running.go | 15 +- core/services/worker/replica_test.go | 28 +- core/services/worker/supervisor.go | 13 +- core/services/worker/worker.go | 8 +- core/services/workerctl/backend.go | 164 +++++++ core/services/workerctl/backend_test.go | 20 + core/services/workerctl/doc.go | 9 + core/services/workerctl/files.go | 64 +++ core/services/workerctl/model.go | 59 +++ core/services/workerctl/progress.go | 29 ++ core/services/workerctl/progress_test.go | 48 ++ core/services/workerctl/wire_test.go | 153 +++++++ .../workerctl/workerctl_suite_test.go | 13 + docs/content/features/distributed-mode.md | 2 + .../distributed/agent_native_executor_test.go | 138 +++--- tests/e2e/distributed/backend_logs_test.go | 128 +++++- .../distributed/distributed_full_flow_test.go | 17 +- tests/e2e/distributed/file_staging_test.go | 2 +- tests/e2e/distributed/foundation_test.go | 10 +- tests/e2e/distributed/job_dispatch_test.go | 9 +- .../e2e/distributed/job_distribution_test.go | 98 +++-- tests/e2e/distributed/managers_test.go | 23 +- tests/e2e/distributed/mcp_nats_test.go | 75 +++- .../distributed/model_config_revision_test.go | 8 +- tests/e2e/distributed/nats_jwt_test.go | 5 +- tests/e2e/distributed/node_lifecycle_test.go | 7 +- .../distributed/prefix_cache_routing_test.go | 2 +- tests/e2e/distributed/router_tracking_test.go | 6 +- tests/e2e/distributed/sse_routes_test.go | 3 +- 155 files changed, 6170 insertions(+), 2057 deletions(-) create mode 100644 .agents/distributed-seams.md create mode 100644 core/application/agent_pool_options_test.go create mode 100644 core/cli/agent_worker_mcp_rpc_test.go create mode 100644 core/cli/agent_worker_mcpci_test.go create mode 100644 core/cli/agent_worker_subject_test.go create mode 100644 core/http/endpoints/mcp/executor_agent_control_test.go create mode 100644 core/services/agentpool/dispatch_chat_test.go create mode 100644 core/services/agents/dispatcher_test.go create mode 100644 core/services/messaging/client_conformance_test.go create mode 100644 core/services/messaging/client_validation_test.go create mode 100644 core/services/messaging/export_test.go create mode 100644 core/services/messaging/interfaces_test.go create mode 100644 core/services/messaging/messagingtest/broadcaster.go create mode 100644 core/services/messaging/subject_rules.go create mode 100644 core/services/messaging/subject_rules_test.go create mode 100644 core/services/messaging/workqueue.go create mode 100644 core/services/messaging/workqueue_nats.go create mode 100644 core/services/messaging/workqueue_nats_test.go create mode 100644 core/services/nodes/agent_rpc_nats.go create mode 100644 core/services/nodes/agent_rpc_nats_test.go create mode 100644 core/services/nodes/client_factory_test.go create mode 100644 core/services/nodes/control_errors_test.go create mode 100644 core/services/nodes/control_nats.go create mode 100644 core/services/nodes/file_stager_dial_test.go create mode 100644 core/services/nodes/noroute_reactions_test.go create mode 100644 core/services/nodes/reconciler_prober_test.go create mode 100644 core/services/nodes/router_liveness_route_test.go create mode 100644 core/services/testutil/export_test.go create mode 100644 core/services/testutil/fakebus_conformance_test.go create mode 100644 core/services/testutil/fakebus_queue_test.go create mode 100644 core/services/testutil/subject_match_test.go create mode 100644 core/services/testutil/testutil_suite_test.go create mode 100644 core/services/worker/control_nats.go create mode 100644 core/services/worker/control_nats_test.go create mode 100644 core/services/worker/control_server.go create mode 100644 core/services/worker/file_staging_verbs_test.go create mode 100644 core/services/workerctl/backend.go create mode 100644 core/services/workerctl/backend_test.go create mode 100644 core/services/workerctl/doc.go create mode 100644 core/services/workerctl/files.go create mode 100644 core/services/workerctl/model.go create mode 100644 core/services/workerctl/progress.go create mode 100644 core/services/workerctl/progress_test.go create mode 100644 core/services/workerctl/wire_test.go create mode 100644 core/services/workerctl/workerctl_suite_test.go 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()