mirror of
https://github.com/mudler/LocalAI.git
synced 2026-10-04 20:14:43 -04:00
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:
1 parent
776d9cc30b
commit
e4fa051ee3
155 files changed
+6170
-2057
No files matched your search
@@ -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.
|
||||
@@ -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)))
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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",
|
||||
}))
|
||||
})
|
||||
})
|
||||
@@ -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"))
|
||||
})
|
||||
})
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
@@ -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.
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
})
|
||||
@@ -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,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),
|
||||
)
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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"))
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
})
|
||||
@@ -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"))
|
||||
})
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
})
|
||||
@@ -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 }
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
})
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
})
|
||||
|
||||
@@ -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"},
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
@@ -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
Reference in new issue
Block a user