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 <mudler@localai.io>

* test(messaging): cover BroadcastRoots, ControlRoots and SubjectRoot

Assisted-by: Claude:claude-sonnet-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* docs(nodes): state which FileStager implementations return ErrNoRoute

Assisted-by: Claude:claude-sonnet-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>

* refactor(nodes): build backend clients through one node-aware seam

Assisted-by: Claude:claude-sonnet-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>

* docs: describe the distributed transport seams

Assisted-by: Claude:claude-sonnet-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>

* docs: correct comments that overclaim after the seams refactor

Assisted-by: Claude:claude-sonnet-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* refactor(nodes): dial backend probes through the client factory

Assisted-by: Claude:claude-sonnet-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>

* refactor(nodes): dial workers' file servers through a per-node dialer

Assisted-by: Claude:claude-sonnet-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

* 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 <mudler@localai.io>

---------

Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
Co-authored-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
mudler-agentandEttore Di Giacinto authored and GitHub committed 2026-10-02 23:58:41 +02:00
1 parent 776d9cc30b
commit e4fa051ee3
155 files changed
+6170 -2057

No files matched your search

+169
View File
@@ -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.<id>.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.
+1
View File
@@ -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
@@ -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)))
})
})
+15 -10
View File
@@ -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
+16 -2
View File
@@ -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
}
+1 -2
View File
@@ -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)
+105 -95
View File
@@ -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)
}
}
+47
View File
@@ -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())
})
})
+139
View File
@@ -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",
}))
})
})
+31
View File
@@ -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"))
})
})
-1
View File
@@ -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
+3 -1
View File
@@ -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 {
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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)
+6 -6
View File
@@ -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,
+18 -11
View File
@@ -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)
}
@@ -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"))
+20 -20
View File
@@ -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 {
@@ -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))
})
})
+25 -29
View File
@@ -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.
+2 -2
View File
@@ -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)
@@ -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 {
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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
+3 -3
View File
@@ -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{
+4 -4
View File
@@ -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)
+4 -4
View File
@@ -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.
+4 -4
View File
@@ -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,
+4 -4
View File
@@ -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{
+2 -2
View File
@@ -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()
}
+9 -17
View File
@@ -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
@@ -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())
})
})
+2 -2
View File
@@ -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
}
+36 -85
View File
@@ -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)
+159
View File
@@ -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),
)
})
+10 -13
View File
@@ -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
}
+40 -36
View File
@@ -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))
})
})
})
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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{
+2 -2
View File
@@ -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
+4 -4
View File
@@ -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
+19 -198
View File
@@ -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 {
+120 -61
View File
@@ -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
}
+11
View File
@@ -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"`
@@ -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.<nodeID>.backend.install.<opID>.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.
@@ -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))
})
})
})
+10 -14
View File
@@ -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.
@@ -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
})
})
@@ -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))
})
}
})
+5
View File
@@ -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
+12 -4
View File
@@ -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)
@@ -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)
)
@@ -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),
)
})
}
+64
View File
@@ -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)
}
@@ -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)
}
})
})
+6 -228
View File
@@ -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"
}
@@ -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))
})
})
+45
View File
@@ -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)
}
+196
View File
@@ -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
}
@@ -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())
})
})
+132
View File
@@ -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)
}
+388
View File
@@ -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"))
})
})
})
@@ -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"))
})
})
@@ -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"))
})
})
+59
View File
@@ -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)
}
+2 -2
View File
@@ -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,
+4
View File
@@ -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.
@@ -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")
})
})
+53 -36
View File
@@ -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)
}
@@ -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"))
})
+10 -55
View File
@@ -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,
@@ -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() {
@@ -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) {
@@ -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")
@@ -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()
+1 -1
View File
@@ -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()
+1 -1
View File
@@ -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 {
@@ -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)
}
}
@@ -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))
})
})
+22 -5
View File
@@ -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 }
}
+15 -16
View File
@@ -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
@@ -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())
})
})
+11 -11
View File
@@ -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())
@@ -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))
})
})
})
@@ -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())
+17 -19
View File
@@ -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"
}
@@ -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),
)
})
+30 -2
View File
@@ -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() {
@@ -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())
+3 -3
View File
@@ -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.
@@ -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
}
+1 -1
View File
@@ -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 {
+2 -3
View File
@@ -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
@@ -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())
})
})
@@ -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()
})
+3 -3
View File
@@ -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()
@@ -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"},
}
})
+3 -3
View File
@@ -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"},
}
})
@@ -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() {
@@ -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"},
}
})
Loaded 100 of 155 files, more files were not shown because too many files have changed in this diff. Show more