endpoint model discovery and proxy integration

This commit is contained in:
Brandon Hopkins
2026-07-26 16:38:01 -07:00
parent a92cdb7dcd
commit d3909e4faf
22 changed files with 3523 additions and 611 deletions
@@ -0,0 +1,404 @@
package grpc
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
rpproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/shared/management/proto"
nbstatus "github.com/netbirdio/netbird/shared/management/status"
)
type modelDiscoveryCall struct {
result *proto.ModelDiscoveryResult
err error
}
func newModelDiscoveryTestServer(t *testing.T) (*ProxyServiceServer, *testProxyController) {
t.Helper()
controller := newTestProxyController()
server := &ProxyServiceServer{
proxyController: controller,
modelDiscoveryPending: make(map[string]*pendingModelDiscovery),
}
return server, controller
}
func addModelDiscoveryTestConnection(
t *testing.T,
server *ProxyServiceServer,
controller *testProxyController,
proxyID, sessionID, cluster string,
accountID *string,
supportsDiscovery bool,
) (*proxyConnection, context.CancelFunc) {
t.Helper()
ctx, cancel := context.WithCancel(context.Background())
ready := make(chan struct{})
close(ready)
conn := &proxyConnection{
proxyID: proxyID,
sessionID: sessionID,
address: cluster,
accountID: accountID,
capabilities: &proto.ProxyCapabilities{
SupportsModelDiscovery: &supportsDiscovery,
},
syncStream: &syncRecordingStream{},
modelDiscoveryChan: make(chan *proto.ModelDiscoveryRequest, modelDiscoveryQueueSize),
ready: ready,
ctx: ctx,
cancel: cancel,
}
server.connectedProxies.Store(proxyID, conn)
require.NoError(t, controller.RegisterProxyToCluster(context.Background(), cluster, proxyID))
t.Cleanup(func() {
cancel()
server.connectedProxies.Delete(proxyID)
_ = controller.UnregisterProxyFromCluster(context.Background(), cluster, proxyID)
})
return conn, cancel
}
func callDiscoverModels(
server *ProxyServiceServer,
ctx context.Context,
accountID, cluster string,
req *proto.ModelDiscoveryRequest,
) <-chan modelDiscoveryCall {
done := make(chan modelDiscoveryCall, 1)
go func() {
result, err := server.DiscoverModels(ctx, accountID, cluster, req)
done <- modelDiscoveryCall{result: result, err: err}
}()
return done
}
func TestDiscoverModels_CorrelatesResultAndCopiesSensitiveRequest(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
callerReq := &proto.ModelDiscoveryRequest{
RequestId: "caller-controlled",
UpstreamUrl: "http://ollama.internal:11434",
AuthHeaderName: "Authorization",
AuthHeaderValue: "Bearer secret",
SkipTlsVerify: true,
OllamaFallback: true,
}
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", callerReq)
wireReq := <-conn.modelDiscoveryChan
assert.NotEmpty(t, wireReq.GetRequestId())
assert.NotEqual(t, callerReq.GetRequestId(), wireReq.GetRequestId(), "management must own correlation IDs")
assert.Equal(t, callerReq.GetUpstreamUrl(), wireReq.GetUpstreamUrl())
assert.Equal(t, callerReq.GetAuthHeaderName(), wireReq.GetAuthHeaderName())
assert.Equal(t, callerReq.GetAuthHeaderValue(), wireReq.GetAuthHeaderValue())
assert.Equal(t, callerReq.GetSkipTlsVerify(), wireReq.GetSkipTlsVerify())
assert.Equal(t, callerReq.GetOllamaFallback(), wireReq.GetOllamaFallback())
assert.Equal(t, "caller-controlled", callerReq.GetRequestId(), "caller request must not be mutated")
want := &proto.ModelDiscoveryResult{
RequestId: wireReq.GetRequestId(),
Models: []*proto.ModelDiscoveryModel{
{Id: "llama3.2:latest", Label: "llama3.2:latest"},
},
Source: "openai_v1_models",
}
server.completeModelDiscovery(conn, want)
got := <-done
require.NoError(t, got.err)
assert.Equal(t, want, got.result)
}
func TestDiscoverModels_IgnoresMismatchedCorrelation(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
wireReq := <-conn.modelDiscoveryChan
server.completeModelDiscovery(conn, &proto.ModelDiscoveryResult{RequestId: "wrong-request"})
select {
case got := <-done:
t.Fatalf("mismatched request completed discovery: %+v", got)
default:
}
wrongSession := *conn
wrongSession.sessionID = "session-b"
server.completeModelDiscovery(&wrongSession, &proto.ModelDiscoveryResult{RequestId: wireReq.GetRequestId()})
select {
case got := <-done:
t.Fatalf("mismatched proxy session completed discovery: %+v", got)
default:
}
server.completeModelDiscovery(conn, &proto.ModelDiscoveryResult{RequestId: wireReq.GetRequestId()})
got := <-done
require.NoError(t, got.err)
require.NotNil(t, got.result)
}
func TestDiscoverModels_ReturnsProxyReportedErrorForManagerContext(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
wireReq := <-conn.modelDiscoveryChan
server.completeModelDiscovery(conn, &proto.ModelDiscoveryResult{
RequestId: wireReq.GetRequestId(),
Error: "upstream returned HTTP 401",
})
got := <-done
require.NoError(t, got.err)
require.NotNil(t, got.result)
assert.Equal(t, "upstream returned HTTP 401", got.result.GetError())
assert.Equal(t, wireReq.GetRequestId(), got.result.GetRequestId())
}
func TestDiscoverModels_FiltersCapabilityAndAccountScope(t *testing.T) {
tests := []struct {
name string
connectionAccount *string
requestAccount string
supported bool
}{
{
name: "capability absent",
requestAccount: "account-a",
supported: false,
},
{
name: "BYOP account mismatch",
connectionAccount: stringPointer("account-b"),
requestAccount: "account-a",
supported: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com",
tt.connectionAccount, tt.supported,
)
result, err := server.DiscoverModels(
context.Background(),
tt.requestAccount,
"cluster.example.com",
&proto.ModelDiscoveryRequest{UpstreamUrl: "http://ollama.internal:11434"},
)
require.Nil(t, result)
requireStatusType(t, err, nbstatus.PreconditionFailed)
})
}
}
func TestDiscoverModels_RejectsStaleClusterMembershipAfterReconnect(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "new-session", "new-cluster.example.com", nil, true,
)
// Simulate the stale membership left behind when this proxy ID previously
// connected to another cluster and that superseded session skipped cleanup.
require.NoError(t, controller.RegisterProxyToCluster(
context.Background(),
"old-cluster.example.com",
"proxy-a",
))
result, err := server.DiscoverModels(
context.Background(),
"account-a",
"old-cluster.example.com",
&proto.ModelDiscoveryRequest{UpstreamUrl: "http://ollama.internal:11434"},
)
require.Nil(t, result)
requireStatusType(t, err, nbstatus.PreconditionFailed)
select {
case request := <-conn.modelDiscoveryChan:
t.Fatalf("sent credentials to connection in a different cluster: %+v", request)
default:
}
}
func TestDiscoverModels_CancelAndDisconnectCleanPendingRequest(t *testing.T) {
t.Run("caller cancellation", func(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
ctx, cancel := context.WithCancel(context.Background())
done := callDiscoverModels(server, ctx, "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
<-conn.modelDiscoveryChan
cancel()
got := <-done
require.Nil(t, got.result)
requireStatusType(t, got.err, nbstatus.PreconditionFailed)
assert.Empty(t, server.modelDiscoveryPending)
})
t.Run("proxy disconnect", func(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
<-conn.modelDiscoveryChan
server.failModelDiscoveries(conn, rpproxy.ErrModelDiscoveryUnavailable)
got := <-done
require.Nil(t, got.result)
requireStatusType(t, got.err, nbstatus.PreconditionFailed)
assert.Empty(t, server.modelDiscoveryPending)
})
}
func TestDiscoverModels_UsesNextCapableProxyWhenFirstQueueIsFull(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
first, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
second, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-b", "session-b", "cluster.example.com", nil, true,
)
first.modelDiscoveryChan = make(chan *proto.ModelDiscoveryRequest, 1)
first.modelDiscoveryChan <- &proto.ModelDiscoveryRequest{RequestId: "occupied"}
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
wireReq := <-second.modelDiscoveryChan
server.completeModelDiscovery(second, &proto.ModelDiscoveryResult{RequestId: wireReq.GetRequestId()})
got := <-done
require.NoError(t, got.err)
require.NotNil(t, got.result)
}
func TestSenderSkipsCanceledQueuedModelDiscovery(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
stream := conn.syncStream.(*syncRecordingStream)
ctx, cancel := context.WithCancel(context.Background())
done := callDiscoverModels(server, ctx, "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
require.Eventually(t, func() bool {
return len(conn.modelDiscoveryChan) == 1
}, time.Second, time.Millisecond)
cancel()
got := <-done
require.Nil(t, got.result)
requireStatusType(t, got.err, nbstatus.PreconditionFailed)
liveRequest := &proto.ModelDiscoveryRequest{RequestId: "live-request"}
livePending := &pendingModelDiscovery{
proxyID: conn.proxyID,
sessionID: conn.sessionID,
done: make(chan modelDiscoveryCompletion, 1),
}
server.registerModelDiscovery(liveRequest.GetRequestId(), livePending)
defer server.removeModelDiscovery(liveRequest.GetRequestId(), livePending)
conn.modelDiscoveryChan <- liveRequest
errCh := make(chan error, 1)
go server.sender(conn, errCh)
require.Eventually(t, func() bool {
stream.mu.Lock()
defer stream.mu.Unlock()
return len(stream.sent) == 1
}, time.Second, time.Millisecond)
stream.mu.Lock()
sent := append([]*proto.SyncMappingsResponse(nil), stream.sent...)
stream.mu.Unlock()
require.Len(t, sent, 1)
assert.Equal(t, "live-request", sent[0].GetModelDiscoveryRequest().GetRequestId())
}
func TestDrainRecv_DispatchesModelDiscoveryResult(t *testing.T) {
server, _ := newModelDiscoveryTestServer(t)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
conn := &proxyConnection{
proxyID: "proxy-a",
sessionID: "session-a",
ctx: ctx,
cancel: cancel,
}
pending := &pendingModelDiscovery{
proxyID: conn.proxyID,
sessionID: conn.sessionID,
done: make(chan modelDiscoveryCompletion, 1),
}
server.registerModelDiscovery("request-a", pending)
stream := &syncRecordingStream{
recvMsgs: []*proto.SyncMappingsRequest{
{
Msg: &proto.SyncMappingsRequest_ModelDiscoveryResult{
ModelDiscoveryResult: &proto.ModelDiscoveryResult{RequestId: "request-a"},
},
},
},
}
errCh := make(chan error, 1)
go server.drainRecv(conn, stream, errCh)
completion := <-pending.done
require.NoError(t, completion.err)
require.NotNil(t, completion.result)
assert.Equal(t, "request-a", completion.result.GetRequestId())
}
func TestSendModelDiscoveryRequest_UsesOutOfBandSyncField(t *testing.T) {
stream := &syncRecordingStream{}
conn := &proxyConnection{syncStream: stream}
req := &proto.ModelDiscoveryRequest{RequestId: "request-a"}
require.NoError(t, conn.sendModelDiscoveryRequest(req))
require.Len(t, stream.sent, 1)
assert.Empty(t, stream.sent[0].GetMapping())
assert.Equal(t, req, stream.sent[0].GetModelDiscoveryRequest())
}
func stringPointer(value string) *string {
return &value
}
func requireStatusType(t *testing.T, err error, want nbstatus.Type) {
t.Helper()
require.Error(t, err)
statusErr, ok := nbstatus.FromError(err)
require.True(t, ok)
assert.Equal(t, want, statusErr.Type())
}
+283 -22
View File
@@ -15,6 +15,7 @@ import (
"net/http"
"net/url"
"os"
"sort"
"strconv"
"strings"
"sync"
@@ -29,6 +30,7 @@ import (
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/peers"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
@@ -36,7 +38,6 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/management/server/users"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
@@ -84,6 +85,10 @@ type ProxyServiceServer struct {
// Map of connected proxies: proxy_id -> proxy connection
connectedProxies sync.Map
// modelDiscoveryPending correlates out-of-band model-discovery results
// received on SyncMappings with the management request waiting for them.
modelDiscoveryMu sync.Mutex
modelDiscoveryPending map[string]*pendingModelDiscovery
// Manager for access logs
accessLogManager accesslogs.Manager
@@ -143,6 +148,8 @@ const defaultProxyTokenTTL = 5 * time.Minute
const defaultSnapshotBatchSize = 500
const modelDiscoveryQueueSize = 16
func snapshotBatchSizeFromEnv() int {
if v := os.Getenv("NB_PROXY_SNAPSHOT_BATCH_SIZE"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
@@ -171,10 +178,26 @@ type proxyConnection struct {
stream proto.ProxyService_GetMappingUpdateServer
// syncStream is set when the proxy connected via SyncMappings.
// When non-nil, the sender goroutine uses this instead of stream.
syncStream proto.ProxyService_SyncMappingsServer
sendChan chan *proto.GetMappingUpdateResponse
ctx context.Context
cancel context.CancelFunc
syncStream proto.ProxyService_SyncMappingsServer
sendChan chan *proto.GetMappingUpdateResponse
modelDiscoveryChan chan *proto.ModelDiscoveryRequest
// ready closes after the initial snapshot completes and the sender
// goroutine is about to start. Discovery never competes with snapshot
// delivery on the stream.
ready chan struct{}
ctx context.Context
cancel context.CancelFunc
}
type modelDiscoveryCompletion struct {
result *proto.ModelDiscoveryResult
err error
}
type pendingModelDiscovery struct {
proxyID string
sessionID string
done chan modelDiscoveryCompletion
}
func enforceAccountScope(ctx context.Context, requestAccountID string) error {
@@ -192,17 +215,18 @@ func enforceAccountScope(ctx context.Context, requestAccountID string) error {
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, pkceStore *PKCEVerifierStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
ctx, cancel := context.WithCancel(context.Background())
s := &ProxyServiceServer{
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
proxyManager: proxyMgr,
tokenChecker: tokenChecker,
snapshotBatchSize: snapshotBatchSizeFromEnv(),
cancel: cancel,
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
proxyManager: proxyMgr,
tokenChecker: tokenChecker,
snapshotBatchSize: snapshotBatchSizeFromEnv(),
modelDiscoveryPending: make(map[string]*pendingModelDiscovery),
cancel: cancel,
}
go s.cleanupStaleProxies(ctx)
return s
@@ -392,6 +416,7 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest
return fmt.Errorf("send snapshot to proxy %s: %w", params.proxyID, err)
}
close(conn.ready)
errChan := make(chan error, 2)
go s.sender(conn, errChan)
@@ -425,9 +450,10 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings
return fmt.Errorf("send snapshot to proxy %s: %w", params.proxyID, err)
}
close(conn.ready)
errChan := make(chan error, 2)
go s.sender(conn, errChan)
go s.drainRecv(stream, errChan)
go s.drainRecv(conn, stream, errChan)
return s.serveProxyConnection(conn, proxyRecord, errChan, true)
}
@@ -496,6 +522,8 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
connSeed.tokenID = tokenID
connSeed.capabilities = params.capabilities
connSeed.sendChan = make(chan *proto.GetMappingUpdateResponse, 100)
connSeed.modelDiscoveryChan = make(chan *proto.ModelDiscoveryRequest, modelDiscoveryQueueSize)
connSeed.ready = make(chan struct{})
connSeed.ctx = connCtx
connSeed.cancel = cancel
@@ -543,10 +571,13 @@ func (s *ProxyServiceServer) supersedePriorConnection(proxyID, newSessionID stri
// cleanupFailedSnapshot removes the connection from the cluster and store
// after a snapshot send failure.
func (s *ProxyServiceServer) cleanupFailedSnapshot(ctx context.Context, conn *proxyConnection) {
s.failModelDiscoveries(conn, proxy.ErrModelDiscoveryUnavailable)
if s.connectedProxies.CompareAndDelete(conn.proxyID, conn) {
if err := s.proxyController.UnregisterProxyFromCluster(context.Background(), conn.address, conn.proxyID); err != nil {
log.WithContext(ctx).Debugf("cleanup after snapshot failure for proxy %s: %v", conn.proxyID, err)
}
} else {
s.unregisterSupersededProxyCluster(ctx, conn)
}
conn.cancel()
if err := s.proxyManager.Disconnect(context.Background(), conn.proxyID, conn.sessionID); err != nil {
@@ -554,15 +585,19 @@ func (s *ProxyServiceServer) cleanupFailedSnapshot(ctx context.Context, conn *pr
}
}
// drainRecv consumes and discards messages from a bidirectional stream.
// The proxy sends an ack for every incremental update; we don't need them
// after the snapshot phase. Recv errors are forwarded to errChan.
func (s *ProxyServiceServer) drainRecv(stream proto.ProxyService_SyncMappingsServer, errChan chan<- error) {
// drainRecv consumes post-snapshot messages from a bidirectional stream.
// Incremental mapping acks need no action; model-discovery results are
// dispatched to the correlated management caller.
func (s *ProxyServiceServer) drainRecv(conn *proxyConnection, stream proto.ProxyService_SyncMappingsServer, errChan chan<- error) {
for {
if _, err := stream.Recv(); err != nil {
msg, err := stream.Recv()
if err != nil {
errChan <- err
return
}
if result := msg.GetModelDiscoveryResult(); result != nil {
s.completeModelDiscovery(conn, result)
}
}
}
@@ -599,7 +634,9 @@ func (s *ProxyServiceServer) serveProxyConnection(conn *proxyConnection, proxyRe
// disconnectProxy removes the connection from cluster and store, unless it
// has already been superseded by a newer connection.
func (s *ProxyServiceServer) disconnectProxy(conn *proxyConnection) {
s.failModelDiscoveries(conn, proxy.ErrModelDiscoveryUnavailable)
if !s.connectedProxies.CompareAndDelete(conn.proxyID, conn) {
s.unregisterSupersededProxyCluster(context.Background(), conn)
log.Infof("Proxy %s session %s: skipping cleanup, superseded by new connection", conn.proxyID, conn.sessionID)
conn.cancel()
return
@@ -616,6 +653,26 @@ func (s *ProxyServiceServer) disconnectProxy(conn *proxyConnection) {
log.Infof("Proxy %s session %s disconnected", conn.proxyID, conn.sessionID)
}
// unregisterSupersededProxyCluster removes membership owned only by an old
// session after the same proxy ID reconnects at a different address. When the
// address is unchanged, the membership is shared with the replacement and
// must remain registered.
func (s *ProxyServiceServer) unregisterSupersededProxyCluster(ctx context.Context, conn *proxyConnection) {
if conn == nil || s.proxyController == nil {
return
}
currentValue, ok := s.connectedProxies.Load(conn.proxyID)
if ok {
current := currentValue.(*proxyConnection)
if current == conn || current.address == conn.address {
return
}
}
if err := s.proxyController.UnregisterProxyFromCluster(ctx, conn.address, conn.proxyID); err != nil {
log.WithContext(ctx).Debugf("cleanup superseded cluster membership for proxy %s: %v", conn.proxyID, err)
}
}
// sendSnapshotSync sends the initial snapshot with back-pressure: it sends
// one batch, then waits for the proxy to ack before sending the next.
func (s *ProxyServiceServer) sendSnapshotSync(ctx context.Context, conn *proxyConnection, stream proto.ProxyService_SyncMappingsServer) error {
@@ -852,6 +909,20 @@ func (s *ProxyServiceServer) sender(conn *proxyConnection, errChan chan<- error)
return
}
log.WithContext(conn.ctx).Tracef("Send response to proxy %s", conn.proxyID)
case req := <-conn.modelDiscoveryChan:
// The API caller may have timed out while this request waited
// behind other stream work. Pending correlation is removed on
// cancellation, so skip stale probes before they reach the
// proxy and expose a stored credential unnecessarily.
if !s.isModelDiscoveryPendingFor(conn, req.GetRequestId()) {
continue
}
if err := conn.sendModelDiscoveryRequest(req); err != nil {
errChan <- err
log.WithContext(conn.ctx).Tracef("Failed to send model discovery request to proxy %s: %v", conn.proxyID, err)
return
}
log.WithContext(conn.ctx).Tracef("Sent model discovery request to proxy %s", conn.proxyID)
case <-conn.ctx.Done():
return
}
@@ -869,6 +940,15 @@ func (conn *proxyConnection) sendResponse(resp *proto.GetMappingUpdateResponse)
return conn.stream.Send(resp)
}
func (conn *proxyConnection) sendModelDiscoveryRequest(req *proto.ModelDiscoveryRequest) error {
if conn.syncStream == nil {
return proxy.ErrModelDiscoveryUnavailable
}
return conn.syncStream.Send(&proto.SyncMappingsResponse{
ModelDiscoveryRequest: req,
})
}
// SendAccessLog processes access log from proxy
func (s *ProxyServiceServer) SendAccessLog(ctx context.Context, req *proto.SendAccessLogRequest) (*proto.SendAccessLogResponse, error) {
accessLog := req.GetLog()
@@ -1017,6 +1097,181 @@ func (s *ProxyServiceServer) GetConnectedProxyURLs() []string {
return urls
}
// DiscoverModels sends one correlated model-discovery request to a capable
// proxy in clusterAddr and waits for its result. Only the stored connection's
// cluster/account scope is used for routing; no routing identity is carried in
// the request delivered to the proxy.
func (s *ProxyServiceServer) DiscoverModels(ctx context.Context, accountID, clusterAddr string, req *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
if req == nil {
return nil, nbstatus.Errorf(nbstatus.InvalidArgument, "model discovery request is required")
}
if strings.TrimSpace(accountID) == "" {
return nil, nbstatus.Errorf(nbstatus.InvalidArgument, "account ID is required for model discovery")
}
if strings.TrimSpace(clusterAddr) == "" {
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "agent network proxy cluster is not configured")
}
if s.proxyController == nil {
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery is unavailable for the configured proxy cluster")
}
proxyIDs := s.proxyController.GetProxiesForCluster(clusterAddr)
sort.Strings(proxyIDs)
request := &proto.ModelDiscoveryRequest{
RequestId: uuid.NewString(),
UpstreamUrl: req.GetUpstreamUrl(),
AuthHeaderName: req.GetAuthHeaderName(),
AuthHeaderValue: req.GetAuthHeaderValue(),
SkipTlsVerify: req.GetSkipTlsVerify(),
OllamaFallback: req.GetOllamaFallback(),
}
for _, proxyID := range proxyIDs {
connVal, ok := s.connectedProxies.Load(proxyID)
if !ok {
continue
}
conn := connVal.(*proxyConnection)
// Cluster membership can briefly retain a proxy ID from a superseded
// session. The live connection is authoritative: never send provider
// credentials unless it is connected to the exact requested cluster.
if conn.address != clusterAddr {
continue
}
if !modelDiscoveryCapable(conn, accountID) {
continue
}
pending := &pendingModelDiscovery{
proxyID: conn.proxyID,
sessionID: conn.sessionID,
done: make(chan modelDiscoveryCompletion, 1),
}
s.registerModelDiscovery(request.GetRequestId(), pending)
queued := false
select {
case conn.modelDiscoveryChan <- request:
queued = true
case <-ctx.Done():
s.removeModelDiscovery(request.GetRequestId(), pending)
return nil, modelDiscoveryContextError(ctx.Err())
default:
s.removeModelDiscovery(request.GetRequestId(), pending)
}
if !queued {
continue
}
defer s.removeModelDiscovery(request.GetRequestId(), pending)
select {
case completion := <-pending.done:
if completion.err != nil {
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery proxy disconnected")
}
return completion.result, nil
case <-ctx.Done():
return nil, modelDiscoveryContextError(ctx.Err())
case <-conn.ctx.Done():
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery proxy disconnected")
}
}
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "no connected proxy in the configured cluster supports model discovery")
}
func modelDiscoveryCapable(conn *proxyConnection, accountID string) bool {
if conn == nil || conn.syncStream == nil || conn.modelDiscoveryChan == nil || conn.ready == nil {
return false
}
if conn.ctx == nil || conn.ctx.Err() != nil {
return false
}
if conn.accountID != nil && *conn.accountID != accountID {
return false
}
if conn.capabilities == nil || !conn.capabilities.GetSupportsModelDiscovery() {
return false
}
select {
case <-conn.ready:
return true
default:
return false
}
}
func modelDiscoveryContextError(err error) error {
if errors.Is(err, context.DeadlineExceeded) {
return nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery timed out")
}
return nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery was canceled")
}
func (s *ProxyServiceServer) registerModelDiscovery(requestID string, pending *pendingModelDiscovery) {
s.modelDiscoveryMu.Lock()
defer s.modelDiscoveryMu.Unlock()
if s.modelDiscoveryPending == nil {
s.modelDiscoveryPending = make(map[string]*pendingModelDiscovery)
}
s.modelDiscoveryPending[requestID] = pending
}
func (s *ProxyServiceServer) removeModelDiscovery(requestID string, pending *pendingModelDiscovery) {
s.modelDiscoveryMu.Lock()
defer s.modelDiscoveryMu.Unlock()
if s.modelDiscoveryPending[requestID] == pending {
delete(s.modelDiscoveryPending, requestID)
}
}
func (s *ProxyServiceServer) isModelDiscoveryPendingFor(conn *proxyConnection, requestID string) bool {
if conn == nil || requestID == "" {
return false
}
s.modelDiscoveryMu.Lock()
defer s.modelDiscoveryMu.Unlock()
pending := s.modelDiscoveryPending[requestID]
return pending != nil && pending.proxyID == conn.proxyID && pending.sessionID == conn.sessionID
}
func (s *ProxyServiceServer) completeModelDiscovery(conn *proxyConnection, result *proto.ModelDiscoveryResult) {
if conn == nil || result == nil || result.GetRequestId() == "" {
return
}
s.modelDiscoveryMu.Lock()
pending := s.modelDiscoveryPending[result.GetRequestId()]
s.modelDiscoveryMu.Unlock()
if pending == nil || pending.proxyID != conn.proxyID || pending.sessionID != conn.sessionID {
return
}
select {
case pending.done <- modelDiscoveryCompletion{result: result}:
default:
}
}
func (s *ProxyServiceServer) failModelDiscoveries(conn *proxyConnection, err error) {
if conn == nil {
return
}
s.modelDiscoveryMu.Lock()
pending := make([]*pendingModelDiscovery, 0)
for _, item := range s.modelDiscoveryPending {
if item.proxyID == conn.proxyID && item.sessionID == conn.sessionID {
pending = append(pending, item)
}
}
s.modelDiscoveryMu.Unlock()
for _, item := range pending {
select {
case item.done <- modelDiscoveryCompletion{err: err}:
default:
}
}
}
// SendServiceUpdateToCluster sends a service update to all proxy servers in a specific cluster.
// If clusterAddr is empty, broadcasts to all connected proxy servers (backward compatibility).
// For create/update operations a unique one-time auth token is generated per
@@ -1049,6 +1304,12 @@ func (s *ProxyServiceServer) SendServiceUpdateToCluster(ctx context.Context, upd
continue
}
conn := connVal.(*proxyConnection)
// Membership can retain this proxy ID from a superseded session in a
// different cluster. The live connection address is authoritative,
// especially because Agent Network mappings can contain credentials.
if conn.address != clusterAddr {
continue
}
if conn.accountID != nil && update.AccountId != "" && *conn.accountID != update.AccountId {
continue
}
@@ -42,6 +42,10 @@ func newTestProxyController() *testProxyController {
func (c *testProxyController) SendServiceUpdateToCluster(_ context.Context, _ string, _ *proto.ProxyMapping, _ string) {
}
func (c *testProxyController) DiscoverModels(_ context.Context, _, _ string, _ *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
return nil, proxy.ErrModelDiscoveryUnavailable
}
func (c *testProxyController) GetOIDCValidationConfig() proxy.OIDCValidationConfig {
return proxy.OIDCValidationConfig{}
}
@@ -218,6 +222,65 @@ func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
assert.Empty(t, msg2.AuthToken)
}
func TestSendServiceUpdateToCluster_RejectsStaleClusterMembership(t *testing.T) {
ctx := context.Background()
s := &ProxyServiceServer{
tokenStore: NewOneTimeTokenStore(ctx, testCacheStore(t)),
}
controller := newTestProxyController()
s.SetProxyController(controller)
ch := registerFakeProxy(s, "proxy-a", "new-cluster.example.com")
require.NoError(t, controller.RegisterProxyToCluster(ctx, "old-cluster.example.com", "proxy-a"))
s.SendServiceUpdateToCluster(ctx, &proto.ProxyMapping{
Type: proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED,
Id: "agent-network-provider-1",
AccountId: "account-1",
Domain: "agent.example.com",
Path: []*proto.PathMapping{
{Path: "/", Target: "http://ollama.internal:11434/"},
},
}, "old-cluster.example.com")
assert.True(t, drainEmpty(ch), "stale membership must not route a mapping to a different live cluster")
}
func TestDisconnectProxy_RemovesSupersededMembershipFromOldCluster(t *testing.T) {
controller := newTestProxyController()
server := &ProxyServiceServer{proxyController: controller}
oldCtx, cancelOld := context.WithCancel(context.Background())
newCtx, cancelNew := context.WithCancel(context.Background())
t.Cleanup(cancelOld)
t.Cleanup(cancelNew)
oldConnection := &proxyConnection{
proxyID: "proxy-a",
sessionID: "old-session",
address: "old-cluster.example.com",
ctx: oldCtx,
cancel: cancelOld,
}
newConnection := &proxyConnection{
proxyID: "proxy-a",
sessionID: "new-session",
address: "new-cluster.example.com",
ctx: newCtx,
cancel: cancelNew,
}
require.NoError(t, controller.RegisterProxyToCluster(context.Background(), oldConnection.address, oldConnection.proxyID))
require.NoError(t, controller.RegisterProxyToCluster(context.Background(), newConnection.address, newConnection.proxyID))
server.connectedProxies.Store(newConnection.proxyID, newConnection)
server.disconnectProxy(oldConnection)
assert.Empty(t, controller.GetProxiesForCluster(oldConnection.address))
assert.Equal(t, []string{"proxy-a"}, controller.GetProxiesForCluster(newConnection.address))
current, ok := server.connectedProxies.Load(newConnection.proxyID)
require.True(t, ok)
assert.Same(t, newConnection, current)
}
func TestSendServiceUpdate_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))