iop/apps/edge/internal/service/run_dispatch.go

1446 lines
46 KiB
Go

package service
import (
"context"
"fmt"
"strings"
"sync"
"time"
"google.golang.org/protobuf/types/known/structpb"
edgenode "iop/apps/edge/internal/node"
"iop/packages/go/config"
eventpkg "iop/packages/go/events"
iop "iop/proto/gen/iop"
)
// contextClassLong is the ContextClass value that marks a request for
// long-context admission (long slot gating and reservation).
const contextClassLong = "long"
type SubmitRunRequest struct {
NodeRef string
RunID string
ModelGroupKey string
Adapter string
Target string
SessionID string
Workspace string
Prompt string
Input map[string]any
Background bool
TimeoutSec int
MaxQueue int
QueueTimeoutMS int
Metadata map[string]string
// EstimatedInputTokens is a conservative token-count approximation of the
// request input. It is used by long-context admission policy to gate
// routing and queueing decisions.
EstimatedInputTokens int
// ContextClass is a classification tag used by long-context admission
// policy. Allowed values are "normal" and "long".
ContextClass string
// ProviderPool signals that this request should be dispatched via the
// provider-pool catalog keyed by ModelGroupKey. Adapter and Target are
// resolved per-candidate by resolveProviderPoolCandidates; the winning
// candidate's ServedTarget is written into Target before BuildRunRequest.
ProviderPool bool
}
// RunDispatch describes a dispatched run in surface-neutral terms. It is the
// metadata any caller (console, HTTP, future RPC) needs after submission.
type RunDispatch struct {
RunID string
NodeID string
NodeLabel string
ModelGroupKey string
Adapter string
Target string
SessionID string
Background bool
TimeoutSec int
EstimatedInputTokens int
ContextClass string
ProviderID string // non-empty for provider-pool dispatches
ProviderType string // non-empty for provider-pool dispatches
ExecutionPath string // non-empty for provider-pool dispatches
QueueReason string
}
// RunStream carries asynchronous events for a dispatched foreground run.
// Background runs leave both channels nil.
type RunStream struct {
Events <-chan *iop.RunEvent
NodeEvents <-chan *iop.EdgeNodeEvent
}
// RunResult describes the interface for a submitted execution run, allowing
// surface-neutral consumption of its dispatch info and event stream.
type RunResult interface {
Dispatch() RunDispatch
Stream() RunStream
Close()
WaitTimeout() time.Duration
}
// RunHandle is the legacy combined surface kept for existing callers; it
// embeds the surface-neutral DTOs so HTTP/API code can consume RunDispatch
// or RunStream directly without depending on this struct.
type RunHandle struct {
RunDispatch
RunStream
closeOnce sync.Once
close func()
}
func (h *RunHandle) Close() {
if h != nil && h.close != nil {
h.closeOnce.Do(h.close)
}
}
func (h *RunHandle) WaitTimeout() time.Duration {
if h == nil {
return time.Duration(DefaultTimeoutSec+5) * time.Second
}
return time.Duration(normalizeTimeoutSec(h.TimeoutSec)+5) * time.Second
}
func (h *RunHandle) Dispatch() RunDispatch {
if h == nil {
return RunDispatch{}
}
return h.RunDispatch
}
func (h *RunHandle) Stream() RunStream {
if h == nil {
return RunStream{}
}
return h.RunStream
}
func (s *Service) SubmitRun(ctx context.Context, req SubmitRunRequest) (RunResult, error) {
if req.ProviderPool && req.ModelGroupKey != "" && s.queue != nil {
return s.submitRunQueued(ctx, req)
}
return s.submitRunDirect(req)
}
func (s *Service) submitRunDirect(req SubmitRunRequest) (RunResult, error) {
entry, err := s.ResolveNode(req.NodeRef)
if err != nil {
return nil, err
}
return s.dispatchToEntry(entry, req)
}
func (s *Service) submitRunQueued(ctx context.Context, req SubmitRunRequest) (RunResult, error) {
candidates, policy, err := s.resolveQueueCandidates(req)
if err != nil {
return nil, err
}
long := req.ContextClass == contextClassLong
selected, queueReason, err := s.queue.admitWithReason(ctx, req.ModelGroupKey, req.Adapter, req.Target, candidates, policy, long)
if err != nil {
return nil, err
}
// A long slot is reserved only for long requests on providers that declare a
// long-context limit; the release/track path must match that reservation.
longReserved := long && selected.longContextCapacity > 0
// Rewrite adapter and target for provider-pool dispatch: the winning candidate
// carries the concrete adapter and served model name determined at selection time.
if selected.adapter != "" {
req.Adapter = selected.adapter
}
if selected.servedTarget != "" {
req.Target = selected.servedTarget
}
runReq, runID, err := BuildRunRequest(req)
if err != nil {
s.queue.releaseSlotWithLong(req.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)
return nil, err
}
// Track inflight before send so the event watcher can release the slot
// even if a terminal event arrives before the Send call completes.
s.queue.trackInflight(req.ModelGroupKey, runID, selected.entry.NodeID, selected.providerID, longReserved)
var runEvents <-chan *iop.RunEvent
var unregisterRun func()
var nodeEvents <-chan *iop.EdgeNodeEvent
var unregisterNode func()
if !runReq.GetBackground() {
if s.events == nil {
s.queue.releaseRun(runID, "no-event-bus")
return nil, fmt.Errorf("event bus is not configured")
}
runEvents, unregisterRun = s.events.SubscribeRun(runID, 4096)
nodeEvents, unregisterNode = s.events.SubscribeNode(selected.entry.NodeID, 16)
}
if err := selected.entry.Client.Send(runReq); err != nil {
s.queue.releaseRun(runID, "send-error")
if unregisterRun != nil {
unregisterRun()
}
if unregisterNode != nil {
unregisterNode()
}
return nil, err
}
return &RunHandle{
RunDispatch: RunDispatch{
RunID: runID,
NodeID: selected.entry.NodeID,
NodeLabel: nodeLabel(selected.entry),
ModelGroupKey: req.ModelGroupKey,
Adapter: runReq.GetAdapter(),
Target: runReq.GetTarget(),
SessionID: runReq.GetSessionId(),
Background: runReq.GetBackground(),
TimeoutSec: int(runReq.GetTimeoutSec()),
EstimatedInputTokens: req.EstimatedInputTokens,
ContextClass: req.ContextClass,
ProviderID: selected.providerID,
ProviderType: selected.providerType,
ExecutionPath: string(selected.executionPath),
QueueReason: queueReason,
},
RunStream: RunStream{
Events: runEvents,
NodeEvents: nodeEvents,
},
close: func() {
if unregisterRun != nil {
unregisterRun()
}
if unregisterNode != nil {
unregisterNode()
}
},
}, nil
}
// resolveQueueCandidates returns candidate nodes and the group policy for the
// given request. Provider-pool requests are resolved via the model catalog;
// legacy requests are filtered by adapter/target capability.
func (s *Service) resolveQueueCandidates(req SubmitRunRequest) ([]candidateNode, groupPolicy, error) {
store, catalog := s.runtimeConfigSnapshot()
if req.ProviderPool {
return s.resolveProviderPoolCandidates(req, store, catalog)
}
if req.NodeRef != "" {
entry, err := s.ResolveNode(req.NodeRef)
if err != nil {
return nil, groupPolicy{}, err
}
cap := defaultNodeCapacity
if store != nil {
if rec, ok := store.FindByID(entry.NodeID); ok {
res := resolveAdapterForNode(rec, req.Adapter, req.Target)
if !res.supported {
msg := fmt.Sprintf("node %q does not support adapter %q target %q", entry.NodeID, req.Adapter, req.Target)
if res.ambiguous {
msg = fmt.Sprintf("node %q: adapter %q is ambiguous (multiple enabled instances); use an instance key", entry.NodeID, req.Adapter)
}
return nil, groupPolicy{}, fmt.Errorf("%s", msg)
}
cap = res.capacity
}
}
policy := groupPolicyFromRequestOrStore(req, store, []*edgenode.NodeEntry{entry})
return []candidateNode{{entry: entry, capacity: cap}}, policy, nil
}
all := s.registry.All()
if len(all) == 0 {
return nil, groupPolicy{}, fmt.Errorf("no nodes connected")
}
var candidates []candidateNode
for _, entry := range all {
cap := defaultNodeCapacity
if store != nil {
rec, ok := store.FindByID(entry.NodeID)
if ok {
res := resolveAdapterForNode(rec, req.Adapter, req.Target)
if !res.supported {
continue
}
cap = res.capacity
}
}
candidates = append(candidates, candidateNode{entry: entry, capacity: cap})
}
if len(candidates) == 0 {
return nil, groupPolicy{}, fmt.Errorf("no nodes support adapter %q target %q", req.Adapter, req.Target)
}
entries := make([]*edgenode.NodeEntry, len(candidates))
for i, c := range candidates {
entries[i] = c.entry
}
policy := groupPolicyFromRequestOrStore(req, store, entries)
return candidates, policy, nil
}
// adapterResolution holds the resolved capacity and queue policy for a single
// node/adapter combination. ambiguous is true when a type-name lookup (e.g.,
// "ollama") matches 2+ enabled instances on that node, mirroring the Node
// router's exact-instance-key/ambiguity contract.
type adapterResolution struct {
supported bool
ambiguous bool
capacity int
maxQueue int
queueTimeoutMS int
}
func positiveOr(v, fallback int) int {
if v > 0 {
return v
}
return fallback
}
// resolveAdapterForNode determines whether a node can handle adapterType/target
// and computes the per-node capacity and queue policy fields.
//
// Resolution order:
// 1. Exact instance Name match across OllamaInstances / VllmInstances /
// OpenAICompatInstances (instance-key route, matching Node router priority).
// 2. Type-name route (e.g. "ollama"): supported only when exactly 1 enabled
// instance of that type exists. 2+ enabled instances → ambiguous (fail,
// same semantics as Node router ambiguity error). 0 instances → fail-open
// for legacy/unconfigured nodes.
// 3. "cli": capability gated by CLI.Enabled and profile name in target.
// 4. Default (unknown adapter type): fail-open.
func resolveAdapterForNode(rec *edgenode.NodeRecord, adapterType, target string) adapterResolution {
if rec == nil {
return adapterResolution{supported: true, capacity: defaultNodeCapacity}
}
concurrencyFallback := positiveOr(rec.Runtime.Concurrency, defaultNodeCapacity)
// Exact instance Name match (highest priority).
for _, inst := range rec.Adapters.OllamaInstances {
if inst.Name == adapterType {
return adapterResolution{
supported: inst.Enabled,
capacity: positiveOr(inst.Capacity, concurrencyFallback),
maxQueue: inst.MaxQueue,
queueTimeoutMS: inst.QueueTimeoutMS,
}
}
}
for _, inst := range rec.Adapters.VllmInstances {
if inst.Name == adapterType {
return adapterResolution{
supported: inst.Enabled,
capacity: positiveOr(inst.Capacity, concurrencyFallback),
maxQueue: inst.MaxQueue,
queueTimeoutMS: inst.QueueTimeoutMS,
}
}
}
for _, inst := range rec.Adapters.OpenAICompatInstances {
if inst.Name == adapterType {
return adapterResolution{
supported: inst.Enabled,
capacity: positiveOr(inst.Capacity, concurrencyFallback),
maxQueue: inst.MaxQueue,
queueTimeoutMS: inst.QueueTimeoutMS,
}
}
}
// Type-name route.
switch adapterType {
case "ollama":
return resolveTypeRoute(rec, ollamaEnabledInstances(rec), concurrencyFallback)
case "vllm":
return resolveTypeRoute(rec, vllmEnabledInstances(rec), concurrencyFallback)
case "openai_compat":
return resolveTypeRoute(rec, openAICompatEnabledInstances(rec), concurrencyFallback)
case "cli":
// CLI has no named multi-instance model; capability is gated by profile.
if !rec.Adapters.CLI.Enabled && len(rec.Adapters.CLI.Profiles) == 0 {
// No CLI config at all → fail-open for legacy/unconfigured nodes.
return adapterResolution{supported: true, capacity: concurrencyFallback}
}
if !rec.Adapters.CLI.Enabled {
return adapterResolution{supported: false}
}
if target == "" {
return adapterResolution{supported: len(rec.Adapters.CLI.Profiles) > 0, capacity: concurrencyFallback}
}
_, ok := rec.Adapters.CLI.Profiles[target]
return adapterResolution{supported: ok, capacity: concurrencyFallback}
default:
return adapterResolution{supported: true, capacity: concurrencyFallback}
}
}
// instanceFields holds the capacity/policy fields extracted from a single adapter instance.
type instanceFields struct {
capacity, maxQueue, queueTimeoutMS int
}
func ollamaEnabledInstances(rec *edgenode.NodeRecord) []instanceFields {
var out []instanceFields
for _, inst := range rec.Adapters.OllamaInstances {
if inst.Enabled {
out = append(out, instanceFields{inst.Capacity, inst.MaxQueue, inst.QueueTimeoutMS})
}
}
return out
}
func vllmEnabledInstances(rec *edgenode.NodeRecord) []instanceFields {
var out []instanceFields
for _, inst := range rec.Adapters.VllmInstances {
if inst.Enabled {
out = append(out, instanceFields{inst.Capacity, inst.MaxQueue, inst.QueueTimeoutMS})
}
}
return out
}
func openAICompatEnabledInstances(rec *edgenode.NodeRecord) []instanceFields {
var out []instanceFields
for _, inst := range rec.Adapters.OpenAICompatInstances {
if inst.Enabled {
out = append(out, instanceFields{inst.Capacity, inst.MaxQueue, inst.QueueTimeoutMS})
}
}
return out
}
// resolveTypeRoute applies Node-router-compatible type-name resolution:
// - 0 enabled instances → fail-open (legacy / unconfigured)
// - 1 enabled instance → use its capacity and policy
// - 2+ enabled instances → ambiguous (reject)
func resolveTypeRoute(rec *edgenode.NodeRecord, enabled []instanceFields, concurrencyFallback int) adapterResolution {
switch len(enabled) {
case 0:
return adapterResolution{supported: true, capacity: concurrencyFallback}
case 1:
f := enabled[0]
return adapterResolution{
supported: true,
capacity: positiveOr(f.capacity, concurrencyFallback),
maxQueue: f.maxQueue,
queueTimeoutMS: f.queueTimeoutMS,
}
default:
return adapterResolution{supported: false, ambiguous: true}
}
}
func groupPolicyFromRequestOrStore(req SubmitRunRequest, store *edgenode.NodeStore, entries []*edgenode.NodeEntry) groupPolicy {
if req.MaxQueue > 0 || req.QueueTimeoutMS > 0 {
maxQueue := req.MaxQueue
if maxQueue <= 0 {
maxQueue = defaultGroupMaxQueue
}
return groupPolicy{
maxQueue: maxQueue,
queueTimeout: time.Duration(req.QueueTimeoutMS) * time.Millisecond,
queueTimeoutSet: true,
}
}
return groupPolicyFromStore(store, entries, req.Adapter, req.Target)
}
// groupPolicyFromStore derives queue policy from the first resolved candidate node.
func groupPolicyFromStore(store *edgenode.NodeStore, entries []*edgenode.NodeEntry, adapterType, target string) groupPolicy {
if store != nil {
for _, e := range entries {
rec, ok := store.FindByID(e.NodeID)
if !ok {
continue
}
res := resolveAdapterForNode(rec, adapterType, target)
if !res.supported {
continue
}
if res.maxQueue > 0 || res.queueTimeoutMS > 0 {
p := groupPolicy{
maxQueue: positiveOr(res.maxQueue, defaultGroupMaxQueue),
queueTimeout: time.Duration(res.queueTimeoutMS) * time.Millisecond,
queueTimeoutSet: true,
}
return p
}
}
}
return groupPolicy{maxQueue: defaultGroupMaxQueue, queueTimeout: defaultQueueTimeout, queueTimeoutSet: true}
}
// providerCanServe checks whether the provider advertises the served model in
// its own models list (defensive SDD compliance).
func providerCanServe(prov config.NodeProviderConf, servedModel string) bool {
for _, m := range prov.Models {
if m == servedModel {
return true
}
}
return false
}
// providerAdapterKey returns the dispatch adapter key for a provider.
// For legacy/compat providers the explicit Adapter field is used; for
// provider-first providers (Adapter is empty) the provider ID is used.
func providerAdapterKey(prov config.NodeProviderConf) string {
if k := strings.TrimSpace(prov.Adapter); k != "" {
return k
}
return prov.ID
}
// isProviderAvailable checks provider health status. Only "available" (and
// optionally "healthy" as an alias) are considered dispatchable.
func isProviderAvailable(health string) bool {
h := strings.ToLower(strings.TrimSpace(health))
return h == "available" || h == "healthy"
}
// isProviderAdapterInstanceValid is a defensive check in provider-pool candidate
// resolution. It returns false only when the adapter name resolves to a disabled
// exact instance, an ambiguous type route (2+ enabled instances of that type),
// a type route with zero enabled instances, or an unknown/missing adapter key.
// Only enabled exact instances and single-enabled type routes are considered valid.
func isProviderAdapterInstanceValid(rec *edgenode.NodeRecord, adapter string) bool {
if rec == nil {
return true
}
for _, inst := range rec.Adapters.OllamaInstances {
if inst.Name == adapter {
return inst.Enabled
}
}
for _, inst := range rec.Adapters.VllmInstances {
if inst.Name == adapter {
return inst.Enabled
}
}
for _, inst := range rec.Adapters.OpenAICompatInstances {
if inst.Name == adapter {
return inst.Enabled
}
}
if adapter == "cli" {
return rec.Adapters.CLI.Enabled
}
switch adapter {
case "ollama":
// Match edgevalidate.buildAdapterIndex: legacy ollama.Enabled counts as one enabled instance/key.
count := 0
if rec.Adapters.Ollama.Enabled {
count++
}
count += len(ollamaEnabledInstances(rec))
return count == 1
case "vllm":
count := 0
if rec.Adapters.Vllm.Enabled {
count++
}
count += len(vllmEnabledInstances(rec))
return count == 1
case "openai_compat":
count := 0
if rec.Adapters.OpenAICompat.Enabled {
count++
}
count += len(openAICompatEnabledInstances(rec))
return count == 1
}
return false // unknown/custom adapter key: excluded from provider-pool candidates
}
// classifyProviderExecutionPath classifies a provider's execution path based on
// its type. OpenAI-compatible aliases (openai_compat, openai_api, vllm, vllm-mlx,
// lemonade, sglang, seulgivibe_claude, seulgivibe_openai) are routed to the
// tunnel/passthrough path. Ollama, CLI, and unknown/native types use the
// normalized path. This mirrors the SDD requirement that OpenAI-compatible
// callers go through passthrough while Ollama/CLI/native use normalized.
func classifyProviderExecutionPath(providerType string) providerExecutionPath {
switch strings.ToLower(strings.TrimSpace(providerType)) {
case "openai_compat", "openai_api", "vllm", "vllm-mlx", "lemonade", "sglang",
"seulgivibe_claude", "seulgivibe_openai":
return providerExecutionPathTunnel
default:
// ollama, cli, and any unknown/native type → normalized.
return providerExecutionPathNormalized
}
}
// resolveProviderPoolCandidates builds candidates for a provider-pool request.
// It scans connected nodes for providers referenced in the model catalog entry
// that matches req.ModelGroupKey, then assembles per-provider candidateNodes
// carrying the concrete served model name for target rewrite after admission.
// Filters: served-model membership, dispatch adapter presence, and available health.
func (s *Service) resolveProviderPoolCandidates(req SubmitRunRequest, store *edgenode.NodeStore, catalog []config.ModelCatalogEntry) ([]candidateNode, groupPolicy, error) {
var catalogEntry *config.ModelCatalogEntry
for i := range catalog {
if catalog[i].ID == req.ModelGroupKey {
catalogEntry = &catalog[i]
break
}
}
if catalogEntry == nil {
return nil, groupPolicy{}, fmt.Errorf("provider pool model %q not found in catalog", req.ModelGroupKey)
}
all := s.registry.All()
if len(all) == 0 {
return nil, groupPolicy{}, fmt.Errorf("no nodes connected")
}
var candidates []candidateNode
var policy groupPolicy
policySet := false
for _, entry := range all {
if store == nil {
continue
}
rec, ok := store.FindByID(entry.NodeID)
if !ok {
continue
}
for _, prov := range rec.Providers {
servedModel, inCatalog := catalogEntry.Providers[prov.ID]
if !inCatalog {
continue
}
// Exclude disabled providers from dispatch.
if !config.ProviderEnabled(prov) {
continue
}
// Defensive SDD compliance: served target must be in provider's own models list.
if !providerCanServe(prov, servedModel) {
continue
}
// Derive dispatch adapter key: explicit adapter wins; provider-first uses provider ID.
adapterKey := providerAdapterKey(prov)
if strings.TrimSpace(prov.Adapter) != "" {
// Legacy/compat: adapter must resolve to an enabled instance on this node.
if !isProviderAdapterInstanceValid(rec, adapterKey) {
continue
}
}
// Only available/healthy providers are dispatchable.
if !isProviderAvailable(prov.Health) {
continue
}
// SDD compliance: capacity 0 or unknown providers are excluded from
// dispatch candidates. Do NOT fall back to runtime/default concurrency.
if prov.Capacity <= 0 {
continue
}
cap := prov.Capacity
candidates = append(candidates, candidateNode{
entry: entry,
capacity: cap,
priority: prov.Priority,
providerID: prov.ID,
providerType: prov.Type,
adapter: adapterKey,
servedTarget: servedModel,
longContextCapacity: prov.LongContextCapacity,
executionPath: classifyProviderExecutionPath(prov.Type),
})
if !policySet && (prov.MaxQueue > 0 || prov.QueueTimeoutMS > 0) {
policy = groupPolicy{
maxQueue: positiveOr(prov.MaxQueue, defaultGroupMaxQueue),
queueTimeout: time.Duration(prov.QueueTimeoutMS) * time.Millisecond,
queueTimeoutSet: true,
}
policySet = true
}
}
}
if len(candidates) == 0 {
return nil, groupPolicy{}, fmt.Errorf("no connected nodes support provider pool model %q", req.ModelGroupKey)
}
if !policySet {
policy = groupPolicy{maxQueue: defaultGroupMaxQueue, queueTimeout: defaultQueueTimeout, queueTimeoutSet: true}
}
return candidates, policy, nil
}
func (s *Service) dispatchToEntry(entry *edgenode.NodeEntry, req SubmitRunRequest) (RunResult, error) {
runReq, runID, err := BuildRunRequest(req)
if err != nil {
return nil, err
}
var runEvents <-chan *iop.RunEvent
var unregisterRun func()
var nodeEvents <-chan *iop.EdgeNodeEvent
var unregisterNode func()
if !runReq.GetBackground() {
if s.events == nil {
return nil, fmt.Errorf("event bus is not configured")
}
runEvents, unregisterRun = s.events.SubscribeRun(runID, 4096)
nodeEvents, unregisterNode = s.events.SubscribeNode(entry.NodeID, 16)
}
if err := entry.Client.Send(runReq); err != nil {
if unregisterRun != nil {
unregisterRun()
}
if unregisterNode != nil {
unregisterNode()
}
return nil, err
}
return &RunHandle{
RunDispatch: RunDispatch{
RunID: runID,
NodeID: entry.NodeID,
NodeLabel: nodeLabel(entry),
ModelGroupKey: req.ModelGroupKey,
Adapter: runReq.GetAdapter(),
Target: runReq.GetTarget(),
SessionID: runReq.GetSessionId(),
Background: runReq.GetBackground(),
TimeoutSec: int(runReq.GetTimeoutSec()),
EstimatedInputTokens: req.EstimatedInputTokens,
ContextClass: req.ContextClass,
QueueReason: "dispatched",
},
RunStream: RunStream{
Events: runEvents,
NodeEvents: nodeEvents,
},
close: func() {
if unregisterRun != nil {
unregisterRun()
}
if unregisterNode != nil {
unregisterNode()
}
},
}, nil
}
// tunnelFrameBuffer is the per-tunnel channel capacity. The router blocks on a
// full buffer instead of dropping so ordered passthrough bytes are never lost
// while the subscriber is alive.
const tunnelFrameBuffer = 4096
// providerTunnelRouter delivers ProviderTunnelFrame messages to the
// request-bound subscriber of their tunnel id. It is intentionally separate
// from events.Bus: tunnel body bytes are the passthrough source of truth and
// must not enter the lossy run event fanout.
type providerTunnelRouter struct {
mu sync.Mutex
subs map[string]*providerTunnelSub
}
type providerTunnelSub struct {
ch chan *iop.ProviderTunnelFrame
done chan struct{}
}
func newProviderTunnelRouter() *providerTunnelRouter {
return &providerTunnelRouter{subs: make(map[string]*providerTunnelSub)}
}
func (r *providerTunnelRouter) subscribe(tunnelID string, buffer int) (<-chan *iop.ProviderTunnelFrame, func()) {
sub := &providerTunnelSub{
ch: make(chan *iop.ProviderTunnelFrame, buffer),
done: make(chan struct{}),
}
r.mu.Lock()
r.subs[tunnelID] = sub
r.mu.Unlock()
var once sync.Once
unsubscribe := func() {
once.Do(func() {
r.mu.Lock()
if cur, ok := r.subs[tunnelID]; ok && cur == sub {
delete(r.subs, tunnelID)
}
r.mu.Unlock()
close(sub.done)
})
}
return sub.ch, unsubscribe
}
// route delivers a frame to its tunnel subscriber. Frames without a subscriber
// are dropped (the request already finished or was cancelled).
func (r *providerTunnelRouter) route(frame *iop.ProviderTunnelFrame) {
if frame == nil {
return
}
r.mu.Lock()
sub := r.subs[frame.GetTunnelId()]
r.mu.Unlock()
if sub == nil {
return
}
select {
case sub.ch <- frame:
case <-sub.done:
}
}
// RouteProviderTunnelFrame delivers a ProviderTunnelFrame received from Edge
// transport to its request-bound tunnel stream. It never publishes to the
// run/node event bus.
func (s *Service) RouteProviderTunnelFrame(frame *iop.ProviderTunnelFrame) {
if s == nil || s.tunnels == nil {
return
}
s.tunnels.route(frame)
}
// SubmitProviderTunnelRequest asks a node to open a raw provider HTTP request
// and relay the response as ordered ProviderTunnelFrame messages. It is the
// passthrough sibling of SubmitRunRequest and shares the provider-pool
// admission gate.
type SubmitProviderTunnelRequest struct {
NodeRef string
RunID string
ModelGroupKey string
Adapter string
Target string
SessionID string
Method string
Path string
Headers map[string]string
Body []byte
// BuildBody, when set, produces the provider request body from the final
// resolved target (provider-pool admission rewrites the target to the
// winning candidate's served model). It takes precedence over Body.
BuildBody func(target string) ([]byte, error)
Stream bool
TimeoutSec int
MaxQueue int
QueueTimeoutMS int
Metadata map[string]string
EstimatedInputTokens int
ContextClass string
ProviderPool bool
}
// ProviderTunnelStream carries the ordered raw provider frames of a dispatched
// tunnel. The channel is closed after the terminal END/ERROR frame or Close.
type ProviderTunnelStream struct {
Frames <-chan *iop.ProviderTunnelFrame
}
// ProviderTunnelResult is the surface-neutral handle for a dispatched provider
// tunnel, mirroring RunResult for the raw passthrough path.
type ProviderTunnelResult interface {
Dispatch() RunDispatch
Stream() ProviderTunnelStream
Close()
WaitTimeout() time.Duration
// SetHeaders allows late-binding of tunnel headers (e.g. provider auth
// validated after provider-pool selection). Implementations may ignore
// this if headers are already baked into the request.
SetHeaders(map[string]string)
}
// ProviderTunnelHandle implements ProviderTunnelResult for tunnels dispatched
// over the Edge-Node socket.
type ProviderTunnelHandle struct {
RunDispatch
TunnelID string
Headers map[string]string
frames <-chan *iop.ProviderTunnelFrame
closeOnce sync.Once
close func()
}
func (h *ProviderTunnelHandle) Dispatch() RunDispatch {
if h == nil {
return RunDispatch{}
}
return h.RunDispatch
}
func (h *ProviderTunnelHandle) Stream() ProviderTunnelStream {
if h == nil {
return ProviderTunnelStream{}
}
return ProviderTunnelStream{Frames: h.frames}
}
func (h *ProviderTunnelHandle) Close() {
if h != nil && h.close != nil {
h.closeOnce.Do(h.close)
}
}
func (h *ProviderTunnelHandle) WaitTimeout() time.Duration {
if h == nil {
return time.Duration(DefaultTimeoutSec+5) * time.Second
}
return time.Duration(normalizeTimeoutSec(h.TimeoutSec)+5) * time.Second
}
func (h *ProviderTunnelHandle) SetHeaders(hdrs map[string]string) {
h.Headers = hdrs
}
// providerPoolPath indicates which execution path was selected for a
// provider-pool dispatch. These constants mirror the candidateNode.executionPath
// values but are exposed at the dispatch surface so callers (e.g. OpenAI
// handler) can distinguish tunnel/passthrough from normalized execution.
type providerPoolPath string
const (
// ProviderPoolPathTunnel means the selected candidate executes via
// provider tunnel / OpenAI-compatible passthrough.
ProviderPoolPathTunnel providerPoolPath = "provider_tunnel"
// ProviderPoolPathNormalized means the selected candidate executes via
// normalized RunEvent path (Ollama/CLI/native).
ProviderPoolPathNormalized providerPoolPath = "normalized"
)
// PrepareTunnel is an optional pre-dispatch hook that lets the caller
// inject or modify headers on a tunnel-path request before buildProviderTunnelRequest
// and the Node Send step run. If PrepareTunnel returns an error, the slot is
// released and no tunnel request is sent. This avoids late header mutation
// after wire dispatch for provider-pool tunnel paths.
type prepareTunnelFunc func(req SubmitProviderTunnelRequest) (SubmitProviderTunnelRequest, error)
// ProviderPoolDispatchRequest bundles the Run and Tunnel surface values for
// a single one-shot provider-pool dispatch. SubmitProviderPool uses exactly
// one queue admission to select a candidate, then dispatches only the
// execution path indicated by the candidate's executionPath.
type ProviderPoolDispatchRequest struct {
Run SubmitRunRequest
Tunnel SubmitProviderTunnelRequest
PrepareTunnel prepareTunnelFunc
}
// ProviderPoolDispatchResult describes which execution path was selected and
// carries the corresponding dispatch result. Exactly one of Run or Tunnel is
// non-nil.
type ProviderPoolDispatchResult struct {
Path providerPoolPath
Run RunResult
Tunnel ProviderTunnelResult
DispatchInfo RunDispatch
}
// SubmitProviderPool is the one-shot provider-pool dispatch surface. It performs
// a single queue admission, selects one candidate from the catalog, and dispatches
// the selected execution path — either tunnel/passthrough (OpenAI-compatible
// providers) or normalized RunEvent (Ollama/CLI/native). This method replaces
// the caller's need to choose between SubmitRun and SubmitProviderTunnel.
func (s *Service) SubmitProviderPool(ctx context.Context, req ProviderPoolDispatchRequest) (*ProviderPoolDispatchResult, error) {
if s.queue == nil {
return nil, fmt.Errorf("model queue is not configured")
}
candidates, policy, err := s.resolveQueueCandidates(req.Run)
if err != nil {
return nil, err
}
long := req.Run.ContextClass == contextClassLong
selected, queueReason, err := s.queue.admitWithReason(ctx, req.Run.ModelGroupKey, req.Run.Adapter, req.Run.Target, candidates, policy, long)
if err != nil {
return nil, err
}
// A long slot is reserved only for long requests on providers that declare a
// long-context limit; the release/track path must match that reservation.
longReserved := long && selected.longContextCapacity > 0
// Rewrite adapter and target for provider-pool dispatch: the winning candidate
// carries the concrete adapter and served model name determined at selection time.
adapter := req.Run.Adapter
if selected.adapter != "" {
adapter = selected.adapter
}
target := req.Run.Target
if selected.servedTarget != "" {
target = selected.servedTarget
}
switch selected.executionPath {
case providerExecutionPathTunnel:
tunnelReq := req.Tunnel
tunnelReq.ProviderPool = true
tunnelReq.ModelGroupKey = req.Run.ModelGroupKey
tunnelReq.Adapter = adapter
tunnelReq.Target = target
tunnelReq.SessionID = req.Run.SessionID
tunnelReq.MaxQueue = req.Run.MaxQueue
tunnelReq.QueueTimeoutMS = req.Run.QueueTimeoutMS
tunnelReq.EstimatedInputTokens = req.Run.EstimatedInputTokens
tunnelReq.ContextClass = req.Run.ContextClass
// Apply pre-dispatch tunnel preparation (e.g. provider auth headers)
// before buildProviderTunnelRequest so headers reach the wire request.
if req.PrepareTunnel != nil {
tunnelReqPrepared, prepErr := req.PrepareTunnel(tunnelReq)
if prepErr != nil {
s.queue.releaseSlotWithLong(req.Run.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)
return nil, prepErr
}
tunnelReq = tunnelReqPrepared
}
tunnelReqResolved, runID, err := buildProviderTunnelRequest(tunnelReq, adapter, target)
if err != nil {
s.queue.releaseSlotWithLong(req.Run.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)
return nil, err
}
s.queue.trackInflight(req.Run.ModelGroupKey, runID, selected.entry.NodeID, selected.providerID, longReserved)
handle, err := s.openProviderTunnel(selected.entry, tunnelReqResolved, tunnelReq, queueReason, true, selected.providerID, selected.providerType, string(selected.executionPath))
if err != nil {
s.queue.releaseRun(runID, "send-error")
return nil, err
}
return &ProviderPoolDispatchResult{
Path: ProviderPoolPathTunnel,
Tunnel: handle,
DispatchInfo: RunDispatch{
RunID: handle.Dispatch().RunID,
NodeID: handle.Dispatch().NodeID,
NodeLabel: handle.Dispatch().NodeLabel,
ModelGroupKey: handle.Dispatch().ModelGroupKey,
Adapter: handle.Dispatch().Adapter,
Target: handle.Dispatch().Target,
SessionID: handle.Dispatch().SessionID,
TimeoutSec: handle.Dispatch().TimeoutSec,
EstimatedInputTokens: handle.Dispatch().EstimatedInputTokens,
ContextClass: handle.Dispatch().ContextClass,
ProviderID: selected.providerID,
ProviderType: selected.providerType,
ExecutionPath: string(selected.executionPath),
QueueReason: queueReason,
},
}, nil
case providerExecutionPathNormalized:
return s.dispatchProviderPoolRun(ctx, req.Run, adapter, target, selected, queueReason, longReserved)
default:
s.queue.releaseSlotWithLong(req.Run.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)
return nil, fmt.Errorf("unknown execution path %q for provider-pool dispatch", selected.executionPath)
}
}
// dispatchProviderPoolRun is a helper that submits a normalized RunRequest after
// provider-pool admission. It tracks inflight, sends the RunRequest to the
// selected entry, and returns a RunResult wrapping the dispatch info.
func (s *Service) dispatchProviderPoolRun(
ctx context.Context,
req SubmitRunRequest,
adapter, target string,
selected *candidateNode,
queueReason string,
longReserved bool,
) (*ProviderPoolDispatchResult, error) {
req.Adapter = adapter
req.Target = target
runReq, runID, err := BuildRunRequest(req)
if err != nil {
s.queue.releaseSlotWithLong(req.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)
return nil, err
}
s.queue.trackInflight(req.ModelGroupKey, runID, selected.entry.NodeID, selected.providerID, longReserved)
var runEvents <-chan *iop.RunEvent
var unregisterRun func()
var nodeEvents <-chan *iop.EdgeNodeEvent
var unregisterNode func()
if !runReq.GetBackground() {
if s.events == nil {
s.queue.releaseRun(runID, "no-event-bus")
return nil, fmt.Errorf("event bus is not configured")
}
runEvents, unregisterRun = s.events.SubscribeRun(runID, 4096)
nodeEvents, unregisterNode = s.events.SubscribeNode(selected.entry.NodeID, 16)
}
if err := selected.entry.Client.Send(runReq); err != nil {
s.queue.releaseRun(runID, "send-error")
if unregisterRun != nil {
unregisterRun()
}
if unregisterNode != nil {
unregisterNode()
}
return nil, err
}
disp := RunDispatch{
RunID: runID,
NodeID: selected.entry.NodeID,
NodeLabel: nodeLabel(selected.entry),
ModelGroupKey: req.ModelGroupKey,
Adapter: runReq.GetAdapter(),
Target: runReq.GetTarget(),
SessionID: runReq.GetSessionId(),
Background: runReq.GetBackground(),
TimeoutSec: int(runReq.GetTimeoutSec()),
EstimatedInputTokens: req.EstimatedInputTokens,
ContextClass: req.ContextClass,
ProviderID: selected.providerID,
ProviderType: selected.providerType,
ExecutionPath: string(selected.executionPath),
QueueReason: queueReason,
}
hr := &RunHandle{
RunDispatch: disp,
RunStream: RunStream{
Events: runEvents,
NodeEvents: nodeEvents,
},
close: func() {
if unregisterRun != nil {
unregisterRun()
}
if unregisterNode != nil {
unregisterNode()
}
},
}
return &ProviderPoolDispatchResult{
Path: ProviderPoolPathNormalized,
Run: hr,
DispatchInfo: disp,
}, nil
}
// SubmitProviderTunnel dispatches a raw provider tunnel request. Provider-pool
// requests go through the same admission gate as SubmitRun; the reserved slot
// is released when the tunnel reaches END/ERROR or the handle is closed
// (cancel), never via the run event bus.
func (s *Service) SubmitProviderTunnel(ctx context.Context, req SubmitProviderTunnelRequest) (ProviderTunnelResult, error) {
if req.ProviderPool && req.ModelGroupKey != "" && s.queue != nil {
return s.submitProviderTunnelQueued(ctx, req)
}
return s.submitProviderTunnelDirect(req)
}
func (s *Service) submitProviderTunnelQueued(ctx context.Context, req SubmitProviderTunnelRequest) (ProviderTunnelResult, error) {
candidates, policy, err := s.resolveQueueCandidates(SubmitRunRequest{
NodeRef: req.NodeRef,
ModelGroupKey: req.ModelGroupKey,
Adapter: req.Adapter,
Target: req.Target,
MaxQueue: req.MaxQueue,
QueueTimeoutMS: req.QueueTimeoutMS,
EstimatedInputTokens: req.EstimatedInputTokens,
ContextClass: req.ContextClass,
ProviderPool: true,
})
if err != nil {
return nil, err
}
long := req.ContextClass == contextClassLong
selected, queueReason, err := s.queue.admitWithReason(ctx, req.ModelGroupKey, req.Adapter, req.Target, candidates, policy, long)
if err != nil {
return nil, err
}
longReserved := long && selected.longContextCapacity > 0
adapter := req.Adapter
if selected.adapter != "" {
adapter = selected.adapter
}
target := req.Target
if selected.servedTarget != "" {
target = selected.servedTarget
}
tunnelReq, runID, err := buildProviderTunnelRequest(req, adapter, target)
if err != nil {
s.queue.releaseSlotWithLong(req.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)
return nil, err
}
// Track inflight before send so END/ERROR/close release paths can find it
// even if the terminal frame arrives before Send returns.
s.queue.trackInflight(req.ModelGroupKey, runID, selected.entry.NodeID, selected.providerID, longReserved)
handle, err := s.openProviderTunnel(selected.entry, tunnelReq, req, queueReason, true, selected.providerID, selected.providerType, string(selected.executionPath))
if err != nil {
s.queue.releaseRun(runID, "send-error")
return nil, err
}
return handle, nil
}
func (s *Service) submitProviderTunnelDirect(req SubmitProviderTunnelRequest) (ProviderTunnelResult, error) {
entry, err := s.ResolveNode(req.NodeRef)
if err != nil {
return nil, err
}
tunnelReq, _, err := buildProviderTunnelRequest(req, req.Adapter, req.Target)
if err != nil {
return nil, err
}
return s.openProviderTunnel(entry, tunnelReq, req, "dispatched", false, "", "", "")
}
// openProviderTunnel subscribes the request-bound frame channel, sends the
// tunnel request to the node, and wraps the raw channel so the reserved
// provider-pool slot is released exactly once on END/ERROR/close.
func (s *Service) openProviderTunnel(entry *edgenode.NodeEntry, tunnelReq *iop.ProviderTunnelRequest, req SubmitProviderTunnelRequest, queueReason string, tracked bool, providerID, providerType, executionPath string) (*ProviderTunnelHandle, error) {
raw, unsubscribe := s.tunnels.subscribe(tunnelReq.GetTunnelId(), tunnelFrameBuffer)
if err := entry.Client.Send(tunnelReq); err != nil {
unsubscribe()
return nil, err
}
runID := tunnelReq.GetRunId()
var releaseOnce sync.Once
release := func(reason string) {
releaseOnce.Do(func() {
if tracked && s.queue != nil {
s.queue.releaseRun(runID, reason)
}
})
}
out := make(chan *iop.ProviderTunnelFrame, tunnelFrameBuffer)
done := make(chan struct{})
go func() {
defer close(out)
for {
select {
case frame, ok := <-raw:
if !ok {
release("tunnel-stream-closed")
return
}
select {
case out <- frame:
case <-done:
release("tunnel-closed")
return
}
if isTerminalProviderTunnelFrame(frame) {
release("tunnel-" + frame.GetKind().String())
return
}
case <-done:
release("tunnel-closed")
return
}
}
}()
return &ProviderTunnelHandle{
RunDispatch: RunDispatch{
RunID: runID,
NodeID: entry.NodeID,
NodeLabel: nodeLabel(entry),
ModelGroupKey: req.ModelGroupKey,
Adapter: tunnelReq.GetAdapter(),
Target: tunnelReq.GetTarget(),
SessionID: tunnelReq.GetSessionId(),
TimeoutSec: int(tunnelReq.GetTimeoutSec()),
EstimatedInputTokens: req.EstimatedInputTokens,
ContextClass: req.ContextClass,
ProviderID: providerID,
ProviderType: providerType,
ExecutionPath: executionPath,
QueueReason: queueReason,
},
TunnelID: tunnelReq.GetTunnelId(),
frames: out,
close: func() {
close(done)
unsubscribe()
release("tunnel-closed")
},
}, nil
}
func isTerminalProviderTunnelFrame(f *iop.ProviderTunnelFrame) bool {
switch f.GetKind() {
case iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END,
iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_ERROR:
return true
default:
return false
}
}
func buildProviderTunnelRequest(req SubmitProviderTunnelRequest, adapter, target string) (*iop.ProviderTunnelRequest, string, error) {
runID := req.RunID
if runID == "" {
runID = NewRunID()
}
body := req.Body
if req.BuildBody != nil {
b, err := req.BuildBody(target)
if err != nil {
return nil, "", err
}
body = b
}
headers := make(map[string]string, len(req.Headers))
for k, v := range req.Headers {
headers[k] = v
}
metadata := make(map[string]string, len(req.Metadata))
for k, v := range req.Metadata {
metadata[k] = v
}
return &iop.ProviderTunnelRequest{
RunId: runID,
TunnelId: runID + "-tunnel",
Adapter: adapter,
Target: target,
Method: req.Method,
Path: req.Path,
Headers: headers,
Body: body,
Stream: req.Stream,
TimeoutSec: int32(normalizeTimeoutSec(req.TimeoutSec)),
Metadata: metadata,
SessionId: NormalizeSessionID(req.SessionID),
}, runID, nil
}
type CancelRunRequest struct {
NodeRef string
RunID string
Adapter string
Target string
SessionID string
}
func BuildCancelRunRequest(req CancelRunRequest) *iop.CancelRequest {
return &iop.CancelRequest{
RunId: req.RunID,
Adapter: req.Adapter,
Target: req.Target,
SessionId: NormalizeSessionID(req.SessionID),
Action: iop.CancelAction_CANCEL_ACTION_CANCEL_RUN,
}
}
func (s *Service) CancelRun(_ context.Context, req CancelRunRequest) (CommandResult, error) {
entry, err := s.ResolveNode(req.NodeRef)
if err != nil {
return CommandResult{}, err
}
cancelReq := BuildCancelRunRequest(req)
if err := entry.Client.Send(cancelReq); err != nil {
return CommandResult{}, err
}
return CommandResult{
NodeID: entry.NodeID,
NodeLabel: nodeLabel(entry),
SessionID: cancelReq.GetSessionId(),
}, nil
}
type TerminateSessionRequest struct {
NodeRef string
Adapter string
Target string
SessionID string
}
// CommandResult is the surface-neutral acknowledgement for one-shot node
// commands (terminate-session, future control RPCs).
type CommandResult struct {
NodeID string
NodeLabel string
SessionID string
}
// TerminateSessionResult is retained as an alias for backward compatibility.
type TerminateSessionResult = CommandResult
func (s *Service) TerminateSession(_ context.Context, req TerminateSessionRequest) (TerminateSessionResult, error) {
entry, err := s.ResolveNode(req.NodeRef)
if err != nil {
return TerminateSessionResult{}, err
}
sessionID := NormalizeSessionID(req.SessionID)
cancelReq := &iop.CancelRequest{
Adapter: req.Adapter,
Target: req.Target,
SessionId: sessionID,
Action: iop.CancelAction_CANCEL_ACTION_TERMINATE_SESSION,
}
if err := entry.Client.Send(cancelReq); err != nil {
return TerminateSessionResult{}, err
}
return TerminateSessionResult{
NodeID: entry.NodeID,
NodeLabel: nodeLabel(entry),
SessionID: sessionID,
}, nil
}
func NewRunID() string {
return fmt.Sprintf("manual-%d", time.Now().UnixNano())
}
func IsNodeDisconnected(event *iop.EdgeNodeEvent) bool {
return event.GetType() == eventpkg.TypeNodeDisconnected
}
func BuildRunRequest(req SubmitRunRequest) (*iop.RunRequest, string, error) {
inputMap := make(map[string]any, len(req.Input)+1)
for k, v := range req.Input {
inputMap[k] = v
}
if req.Prompt != "" {
if _, ok := inputMap["prompt"]; !ok {
inputMap["prompt"] = req.Prompt
}
}
input, err := structpb.NewStruct(inputMap)
if err != nil {
return nil, "", err
}
runID := req.RunID
if runID == "" {
runID = NewRunID()
}
metadata := make(map[string]string, len(req.Metadata))
for k, v := range req.Metadata {
metadata[k] = v
}
return &iop.RunRequest{
RunId: runID,
Adapter: req.Adapter,
Target: req.Target,
Workspace: req.Workspace,
SessionId: NormalizeSessionID(req.SessionID),
SessionMode: iop.RunSessionMode_RUN_SESSION_MODE_CREATE_IF_MISSING,
Background: req.Background,
Input: input,
TimeoutSec: int32(normalizeTimeoutSec(req.TimeoutSec)),
Metadata: metadata,
}, runID, nil
}