mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-07 22:19:08 +02:00
endpoint model discovery and proxy integration
This commit is contained in:
@@ -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())
|
||||
}
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user