mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-03 12:09:09 +02:00
endpoint model discovery and proxy integration
This commit is contained in:
@@ -101,6 +101,17 @@ type Provider struct {
|
||||
// upstream provider + credentials on Portkey's hosted side).
|
||||
ExtraHeaders []ExtraHeader
|
||||
Models []Model
|
||||
// ModelDiscovery configures provider-specific discovery behavior. A nil
|
||||
// profile means discovery is unsupported. The profile never carries
|
||||
// caller-supplied paths: the proxy owns the fixed endpoint allowlist.
|
||||
ModelDiscovery *ModelDiscovery
|
||||
}
|
||||
|
||||
// ModelDiscovery describes the safe, catalog-owned fallback behavior for a
|
||||
// discoverable provider. Every discovery starts with OpenAI-compatible
|
||||
// /v1/models; OllamaFallback permits /api/tags only when that route is absent.
|
||||
type ModelDiscovery struct {
|
||||
OllamaFallback bool
|
||||
}
|
||||
|
||||
// ExtraHeader names a single optional per-provider routing/config
|
||||
@@ -732,6 +743,9 @@ var providers = []Provider{
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#000000",
|
||||
Models: []Model{},
|
||||
ModelDiscovery: &ModelDiscovery{
|
||||
OllamaFallback: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "custom",
|
||||
@@ -815,16 +829,17 @@ func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider {
|
||||
kind = api.AgentNetworkCatalogProviderKindCustom
|
||||
}
|
||||
resp := api.AgentNetworkCatalogProvider{
|
||||
Id: p.ID,
|
||||
Name: p.Name,
|
||||
Description: p.Description,
|
||||
DefaultHost: p.DefaultHost,
|
||||
Kind: kind,
|
||||
AuthMode: api.AgentNetworkCatalogProviderAuthMode(p.EffectiveAuthMode()),
|
||||
AuthHeaderTemplate: p.AuthHeaderTemplate,
|
||||
DefaultContentType: p.DefaultContentType,
|
||||
BrandColor: p.BrandColor,
|
||||
Models: models,
|
||||
Id: p.ID,
|
||||
Name: p.Name,
|
||||
Description: p.Description,
|
||||
DefaultHost: p.DefaultHost,
|
||||
Kind: kind,
|
||||
AuthMode: api.AgentNetworkCatalogProviderAuthMode(p.EffectiveAuthMode()),
|
||||
AuthHeaderTemplate: p.AuthHeaderTemplate,
|
||||
DefaultContentType: p.DefaultContentType,
|
||||
BrandColor: p.BrandColor,
|
||||
Models: models,
|
||||
SupportsModelDiscovery: p.ModelDiscovery != nil,
|
||||
}
|
||||
if len(p.ExtraHeaders) > 0 {
|
||||
extras := make([]api.AgentNetworkCatalogExtraHeader, 0, len(p.ExtraHeaders))
|
||||
|
||||
@@ -23,15 +23,26 @@ func TestOllamaCatalogEntry(t *testing.T) {
|
||||
assert.Equal(t, "application/json", entry.DefaultContentType)
|
||||
assert.Empty(t, entry.ParserID, "Ollama preserves the untagged vLLM/custom routing behavior")
|
||||
assert.Empty(t, entry.Models, "Ollama models are installed dynamically on the configured endpoint")
|
||||
require.NotNil(t, entry.ModelDiscovery)
|
||||
assert.True(t, entry.ModelDiscovery.OllamaFallback)
|
||||
|
||||
wire := entry.ToAPIResponse()
|
||||
assert.Equal(t, "ollama", wire.Id)
|
||||
assert.Equal(t, api.AgentNetworkCatalogProviderKindCustom, wire.Kind)
|
||||
assert.Equal(t, api.AgentNetworkCatalogProviderAuthModeOptional, wire.AuthMode)
|
||||
assert.True(t, wire.SupportsModelDiscovery)
|
||||
assert.NotNil(t, wire.Models)
|
||||
assert.Empty(t, wire.Models)
|
||||
}
|
||||
|
||||
func TestOnlyOllamaSupportsModelDiscovery(t *testing.T) {
|
||||
for _, entry := range All() {
|
||||
supportsDiscovery := entry.ModelDiscovery != nil
|
||||
assert.Equal(t, entry.ID == "ollama", supportsDiscovery, entry.ID)
|
||||
assert.Equal(t, supportsDiscovery, entry.ToAPIResponse().SupportsModelDiscovery, entry.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCatalogAuthenticationModes(t *testing.T) {
|
||||
openAI, ok := Lookup("openai_api")
|
||||
require.True(t, ok)
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/shared/auth"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
type discoveryManagerStub struct {
|
||||
agentnetwork.Manager
|
||||
result *agentNetworkTypes.ModelDiscoveryResult
|
||||
err error
|
||||
accountID string
|
||||
userID string
|
||||
providerID string
|
||||
}
|
||||
|
||||
func (m *discoveryManagerStub) DiscoverProviderModels(_ context.Context, accountID, userID, providerID string) (*agentNetworkTypes.ModelDiscoveryResult, error) {
|
||||
m.accountID = accountID
|
||||
m.userID = userID
|
||||
m.providerID = providerID
|
||||
return m.result, m.err
|
||||
}
|
||||
|
||||
func TestDiscoverProviderModelsHandler(t *testing.T) {
|
||||
manager := &discoveryManagerStub{
|
||||
Manager: agentnetwork.NewManagerMock(),
|
||||
result: &agentNetworkTypes.ModelDiscoveryResult{
|
||||
RequestID: "probe-123",
|
||||
Source: "ollama_api_tags",
|
||||
ProxyCluster: "private.example.com",
|
||||
Models: []agentNetworkTypes.DiscoveredModel{
|
||||
{ID: "llama3.2:latest", Label: "llama3.2:latest"},
|
||||
},
|
||||
},
|
||||
}
|
||||
router := mux.NewRouter()
|
||||
RegisterEndpoints(manager, router)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/agent-network/providers/provider-1/discover-models", nil)
|
||||
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{
|
||||
UserId: testUserID,
|
||||
AccountId: testAccountID,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
|
||||
assert.Equal(t, "no-store", rec.Header().Get("Cache-Control"))
|
||||
var response api.AgentNetworkModelDiscoveryResponse
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &response))
|
||||
assert.Equal(t, testAccountID, manager.accountID)
|
||||
assert.Equal(t, testUserID, manager.userID)
|
||||
assert.Equal(t, "provider-1", manager.providerID)
|
||||
assert.Equal(t, "probe-123", response.RequestId)
|
||||
assert.Equal(t, "private.example.com", response.ProxyCluster)
|
||||
assert.Equal(t, api.AgentNetworkModelDiscoveryResponseSourceOllamaApiTags, response.Source)
|
||||
require.Len(t, response.Models, 1)
|
||||
assert.Equal(t, "llama3.2:latest", response.Models[0].Id)
|
||||
}
|
||||
@@ -35,6 +35,7 @@ func RegisterEndpoints(manager agentnetwork.Manager, router *mux.Router) {
|
||||
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/providers/{providerId}", h.updateProvider).Methods("PUT", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/providers/{providerId}", h.deleteProvider).Methods("DELETE", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/providers/{providerId}/discover-models", h.discoverProviderModels).Methods("POST", "OPTIONS")
|
||||
h.addPolicyEndpoints(router)
|
||||
h.addGuardrailEndpoints(router)
|
||||
h.addSettingsEndpoints(router)
|
||||
@@ -98,6 +99,41 @@ func (h *handler) getProvider(w http.ResponseWriter, r *http.Request) {
|
||||
util.WriteJSONObject(r.Context(), w, provider.ToAPIResponse())
|
||||
}
|
||||
|
||||
func (h *handler) discoverProviderModels(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
providerID := strings.TrimSpace(mux.Vars(r)["providerId"])
|
||||
if providerID == "" {
|
||||
util.WriteError(r.Context(), status.Errorf(status.InvalidArgument, "provider ID is required"), w)
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.manager.DiscoverProviderModels(r.Context(), userAuth.AccountId, userAuth.UserId, providerID)
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
models := make([]api.AgentNetworkDiscoveredModel, 0, len(result.Models))
|
||||
for _, model := range result.Models {
|
||||
models = append(models, api.AgentNetworkDiscoveredModel{
|
||||
Id: model.ID,
|
||||
Label: model.Label,
|
||||
})
|
||||
}
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
util.WriteJSONObject(r.Context(), w, api.AgentNetworkModelDiscoveryResponse{
|
||||
Models: models,
|
||||
Source: api.AgentNetworkModelDiscoveryResponseSource(result.Source),
|
||||
ProxyCluster: result.ProxyCluster,
|
||||
RequestId: result.RequestID,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *handler) createProvider(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
|
||||
@@ -9,6 +9,8 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
@@ -49,6 +51,7 @@ func ensureSessionKeys(p *types.Provider) error {
|
||||
type Manager interface {
|
||||
GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error)
|
||||
GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error)
|
||||
DiscoverProviderModels(ctx context.Context, accountID, userID, providerID string) (*types.ModelDiscoveryResult, error)
|
||||
CreateProvider(ctx context.Context, userID string, provider *types.Provider, bootstrapCluster string) (*types.Provider, error)
|
||||
UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
|
||||
DeleteProvider(ctx context.Context, accountID, userID, providerID string) error
|
||||
@@ -129,8 +132,27 @@ type managerImpl struct {
|
||||
// state; concurrent provider creates would otherwise race.
|
||||
labelRngMu sync.Mutex
|
||||
labelRng *rand.Rand
|
||||
|
||||
// discoveryAttempts provides lightweight per-provider admission control for
|
||||
// the explicit endpoint probe. It prevents duplicate clicks from occupying
|
||||
// multiple proxy control-stream requests at once and adds a short cooldown
|
||||
// after each attempt.
|
||||
discoveryMu sync.Mutex
|
||||
discoveryAttempts map[string]modelDiscoveryAttempt
|
||||
}
|
||||
|
||||
type modelDiscoveryAttempt struct {
|
||||
inFlight bool
|
||||
lastStarted time.Time
|
||||
}
|
||||
|
||||
const (
|
||||
modelDiscoveryTimeout = 10 * time.Second
|
||||
modelDiscoveryCooldown = 2 * time.Second
|
||||
maxDiscoveredModels = 500
|
||||
maxDiscoveredModelLen = 512
|
||||
)
|
||||
|
||||
// NewManager constructs the persistent Agent Network manager. The
|
||||
// manager persists provider/policy/guardrail configuration and, on
|
||||
// every mutation, reconciles the in-memory synthesised reverse-proxy
|
||||
@@ -149,6 +171,7 @@ func NewManager(
|
||||
proxyController: proxyController,
|
||||
reconcileCache: make(map[string]map[string]*proto.ProxyMapping),
|
||||
labelRng: rand.New(rand.NewSource(time.Now().UnixNano())),
|
||||
discoveryAttempts: make(map[string]modelDiscoveryAttempt),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -166,6 +189,202 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid
|
||||
return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
|
||||
}
|
||||
|
||||
// DiscoverProviderModels asks a capable proxy in the account's selected
|
||||
// cluster to query the persisted provider endpoint. The browser supplies only
|
||||
// the provider id: URL, TLS policy, and credential are loaded here so a caller
|
||||
// cannot turn this operation into an arbitrary network probe.
|
||||
func (m *managerImpl) DiscoverProviderModels(ctx context.Context, accountID, userID, providerID string) (*types.ModelDiscoveryResult, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Update); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
providerID = strings.TrimSpace(providerID)
|
||||
if providerID == "" {
|
||||
return nil, status.Errorf(status.InvalidArgument, "provider ID is required")
|
||||
}
|
||||
|
||||
provider, err := m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
entry, ok := catalog.Lookup(provider.ProviderID)
|
||||
if !ok {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "provider references an unknown catalog provider")
|
||||
}
|
||||
if entry.ModelDiscovery == nil {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "provider type %q does not support model discovery", provider.ProviderID)
|
||||
}
|
||||
if m.proxyController == nil {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "model discovery is unavailable")
|
||||
}
|
||||
|
||||
settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
var statusErr *status.Error
|
||||
if errors.As(err, &statusErr) && statusErr.Type() == status.NotFound {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "configure an Agent Network proxy cluster before discovering models")
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
cluster := strings.TrimSpace(settings.Cluster)
|
||||
if cluster == "" {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "configure an Agent Network proxy cluster before discovering models")
|
||||
}
|
||||
|
||||
authHeaderName, authHeaderValue, gcpKey, err := providerAuthHeader(provider)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "provider authentication is not configured correctly")
|
||||
}
|
||||
if gcpKey != "" {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "provider authentication is not supported for model discovery")
|
||||
}
|
||||
|
||||
attemptKey := accountID + "\x00" + provider.ID
|
||||
if !m.beginModelDiscovery(attemptKey, time.Now()) {
|
||||
return nil, status.Errorf(status.TooManyRequests, "model discovery is already running or was requested too recently")
|
||||
}
|
||||
defer m.finishModelDiscovery(attemptKey)
|
||||
|
||||
probeCtx, cancel := context.WithTimeout(ctx, modelDiscoveryTimeout)
|
||||
defer cancel()
|
||||
probeResult, err := m.proxyController.DiscoverModels(probeCtx, accountID, cluster, &proto.ModelDiscoveryRequest{
|
||||
UpstreamUrl: strings.TrimSpace(provider.UpstreamURL),
|
||||
AuthHeaderName: authHeaderName,
|
||||
AuthHeaderValue: authHeaderValue,
|
||||
SkipTlsVerify: provider.SkipTLSVerification,
|
||||
OllamaFallback: entry.ModelDiscovery.OllamaFallback,
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, context.DeadlineExceeded) || errors.Is(probeCtx.Err(), context.DeadlineExceeded) {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "model discovery timed out")
|
||||
}
|
||||
if _, ok := status.FromError(err); ok {
|
||||
return nil, err
|
||||
}
|
||||
return nil, status.Errorf(status.PreconditionFailed, "model discovery could not be completed")
|
||||
}
|
||||
if probeResult == nil {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "proxy returned an empty model discovery response")
|
||||
}
|
||||
if strings.TrimSpace(probeResult.Error) != "" {
|
||||
message := safeModelDiscoveryText(probeResult.Error, 256)
|
||||
if message == "" {
|
||||
message = "proxy reported a discovery failure"
|
||||
}
|
||||
requestID := safeModelDiscoveryText(probeResult.RequestId, 128)
|
||||
if requestID != "" {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "model discovery failed: %s (request_id: %s)", message, requestID)
|
||||
}
|
||||
return nil, status.Errorf(status.PreconditionFailed, "model discovery failed: %s", message)
|
||||
}
|
||||
|
||||
requestID := safeModelDiscoveryText(probeResult.RequestId, 128)
|
||||
if requestID == "" {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "proxy returned an uncorrelated model discovery response")
|
||||
}
|
||||
source := strings.TrimSpace(probeResult.Source)
|
||||
switch source {
|
||||
case "openai_v1_models", "ollama_api_tags":
|
||||
default:
|
||||
return nil, status.Errorf(status.PreconditionFailed, "proxy returned an unsupported model discovery response")
|
||||
}
|
||||
|
||||
return &types.ModelDiscoveryResult{
|
||||
RequestID: requestID,
|
||||
Source: source,
|
||||
ProxyCluster: cluster,
|
||||
Models: normalizeDiscoveredModels(probeResult.Models),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (m *managerImpl) beginModelDiscovery(key string, now time.Time) bool {
|
||||
m.discoveryMu.Lock()
|
||||
defer m.discoveryMu.Unlock()
|
||||
if m.discoveryAttempts == nil {
|
||||
m.discoveryAttempts = make(map[string]modelDiscoveryAttempt)
|
||||
}
|
||||
attempt := m.discoveryAttempts[key]
|
||||
if attempt.inFlight || (!attempt.lastStarted.IsZero() && now.Sub(attempt.lastStarted) < modelDiscoveryCooldown) {
|
||||
return false
|
||||
}
|
||||
attempt.inFlight = true
|
||||
attempt.lastStarted = now
|
||||
m.discoveryAttempts[key] = attempt
|
||||
return true
|
||||
}
|
||||
|
||||
func (m *managerImpl) finishModelDiscovery(key string) {
|
||||
m.discoveryMu.Lock()
|
||||
attempt, ok := m.discoveryAttempts[key]
|
||||
if !ok {
|
||||
m.discoveryMu.Unlock()
|
||||
return
|
||||
}
|
||||
attempt.inFlight = false
|
||||
m.discoveryAttempts[key] = attempt
|
||||
m.discoveryMu.Unlock()
|
||||
|
||||
remaining := time.Until(attempt.lastStarted.Add(modelDiscoveryCooldown))
|
||||
if remaining <= 0 {
|
||||
m.expireModelDiscoveryAttempt(key, attempt.lastStarted)
|
||||
return
|
||||
}
|
||||
time.AfterFunc(remaining, func() {
|
||||
m.expireModelDiscoveryAttempt(key, attempt.lastStarted)
|
||||
})
|
||||
}
|
||||
|
||||
func (m *managerImpl) expireModelDiscoveryAttempt(key string, lastStarted time.Time) {
|
||||
m.discoveryMu.Lock()
|
||||
defer m.discoveryMu.Unlock()
|
||||
attempt, ok := m.discoveryAttempts[key]
|
||||
if ok && !attempt.inFlight && attempt.lastStarted.Equal(lastStarted) {
|
||||
delete(m.discoveryAttempts, key)
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeDiscoveredModels(models []*proto.ModelDiscoveryModel) []types.DiscoveredModel {
|
||||
out := make([]types.DiscoveredModel, 0, min(len(models), maxDiscoveredModels))
|
||||
seen := make(map[string]struct{}, min(len(models), maxDiscoveredModels))
|
||||
for _, model := range models {
|
||||
if model == nil {
|
||||
continue
|
||||
}
|
||||
id := safeModelDiscoveryText(model.Id, maxDiscoveredModelLen)
|
||||
if id == "" || len(id) > maxDiscoveredModelLen {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
label := safeModelDiscoveryText(model.Label, maxDiscoveredModelLen)
|
||||
if label == "" || len(label) > maxDiscoveredModelLen {
|
||||
label = id
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
out = append(out, types.DiscoveredModel{ID: id, Label: label})
|
||||
if len(out) == maxDiscoveredModels {
|
||||
break
|
||||
}
|
||||
}
|
||||
slices.SortFunc(out, func(a, b types.DiscoveredModel) int {
|
||||
return strings.Compare(a.ID, b.ID)
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
func safeModelDiscoveryText(value string, maxLen int) string {
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" || len(value) > maxLen || !utf8.ValidString(value) {
|
||||
return ""
|
||||
}
|
||||
for _, r := range value {
|
||||
if unicode.IsControl(r) {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// CreateProvider persists a new provider for the account. bootstrapCluster
|
||||
// is used only when the per-account agent-network Settings row hasn't
|
||||
// been created yet; otherwise it is ignored (the cluster is pinned on
|
||||
@@ -847,6 +1066,10 @@ func (*mockManager) GetProvider(_ context.Context, _, _, _ string) (*types.Provi
|
||||
return &types.Provider{}, nil
|
||||
}
|
||||
|
||||
func (*mockManager) DiscoverProviderModels(_ context.Context, _, _, _ string) (*types.ModelDiscoveryResult, error) {
|
||||
return nil, status.Errorf(status.PreconditionFailed, "model discovery is unavailable")
|
||||
}
|
||||
|
||||
func (*mockManager) CreateProvider(_ context.Context, _ string, p *types.Provider, _ string) (*types.Provider, error) {
|
||||
return p, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
func newModelDiscoveryManager(t *testing.T) (*managerImpl, *store.MockStore, *permissions.MockManager, *proxy.MockController) {
|
||||
t.Helper()
|
||||
ctrl := gomock.NewController(t)
|
||||
mockStore := store.NewMockStore(ctrl)
|
||||
mockPermissions := permissions.NewMockManager(ctrl)
|
||||
mockProxy := proxy.NewMockController(ctrl)
|
||||
return &managerImpl{
|
||||
store: mockStore,
|
||||
permissionsManager: mockPermissions,
|
||||
proxyController: mockProxy,
|
||||
discoveryAttempts: make(map[string]modelDiscoveryAttempt),
|
||||
}, mockStore, mockPermissions, mockProxy
|
||||
}
|
||||
|
||||
func allowModelDiscovery(mockPermissions *permissions.MockManager) {
|
||||
mockPermissions.EXPECT().
|
||||
ValidateUserPermissions(gomock.Any(), "account-1", "user-1", modules.AgentNetwork, operations.Update).
|
||||
Return(true, context.Background(), nil)
|
||||
}
|
||||
|
||||
func TestDiscoverProviderModelsUsesPersistedProviderAndCluster(t *testing.T) {
|
||||
manager, mockStore, mockPermissions, mockProxy := newModelDiscoveryManager(t)
|
||||
allowModelDiscovery(mockPermissions)
|
||||
|
||||
provider := &types.Provider{
|
||||
ID: "provider-1",
|
||||
AccountID: "account-1",
|
||||
ProviderID: "ollama",
|
||||
UpstreamURL: "http://ollama.internal:11434/base",
|
||||
APIKey: "stored-key",
|
||||
SkipTLSVerification: true,
|
||||
}
|
||||
mockStore.EXPECT().
|
||||
GetAgentNetworkProviderByID(gomock.Any(), store.LockingStrengthNone, "account-1", "provider-1").
|
||||
Return(provider, nil)
|
||||
mockStore.EXPECT().
|
||||
GetAgentNetworkSettings(gomock.Any(), store.LockingStrengthNone, "account-1").
|
||||
Return(&types.Settings{AccountID: "account-1", Cluster: "private.example.com"}, nil)
|
||||
mockProxy.EXPECT().
|
||||
DiscoverModels(gomock.Any(), "account-1", "private.example.com", gomock.Any()).
|
||||
DoAndReturn(func(_ context.Context, _, _ string, request *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
|
||||
assert.Empty(t, request.RequestId, "the control-plane server owns request IDs")
|
||||
assert.Equal(t, provider.UpstreamURL, request.UpstreamUrl)
|
||||
assert.Equal(t, "Authorization", request.AuthHeaderName)
|
||||
assert.Equal(t, "Bearer stored-key", request.AuthHeaderValue)
|
||||
assert.True(t, request.SkipTlsVerify)
|
||||
assert.True(t, request.OllamaFallback)
|
||||
return &proto.ModelDiscoveryResult{
|
||||
RequestId: "probe-123",
|
||||
Source: "openai_v1_models",
|
||||
Models: []*proto.ModelDiscoveryModel{
|
||||
{Id: "zeta", Label: ""},
|
||||
{Id: "alpha", Label: "Alpha"},
|
||||
{Id: "zeta", Label: "duplicate"},
|
||||
{Id: "bad\x00model", Label: "bad"},
|
||||
{Id: strings.Repeat("x", maxDiscoveredModelLen+1), Label: "too long"},
|
||||
},
|
||||
}, nil
|
||||
})
|
||||
|
||||
result, err := manager.DiscoverProviderModels(context.Background(), "account-1", "user-1", "provider-1")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "probe-123", result.RequestID)
|
||||
assert.Equal(t, "openai_v1_models", result.Source)
|
||||
assert.Equal(t, "private.example.com", result.ProxyCluster)
|
||||
assert.Equal(t, []types.DiscoveredModel{
|
||||
{ID: "alpha", Label: "Alpha"},
|
||||
{ID: "zeta", Label: "zeta"},
|
||||
}, result.Models)
|
||||
}
|
||||
|
||||
func TestDiscoverProviderModelsRejectsUnsupportedProvider(t *testing.T) {
|
||||
manager, mockStore, mockPermissions, _ := newModelDiscoveryManager(t)
|
||||
allowModelDiscovery(mockPermissions)
|
||||
mockStore.EXPECT().
|
||||
GetAgentNetworkProviderByID(gomock.Any(), store.LockingStrengthNone, "account-1", "provider-1").
|
||||
Return(&types.Provider{
|
||||
ID: "provider-1",
|
||||
AccountID: "account-1",
|
||||
ProviderID: "openai_api",
|
||||
}, nil)
|
||||
|
||||
_, err := manager.DiscoverProviderModels(context.Background(), "account-1", "user-1", "provider-1")
|
||||
require.Error(t, err)
|
||||
statusErr, ok := status.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, status.PreconditionFailed, statusErr.Type())
|
||||
assert.Contains(t, err.Error(), "does not support")
|
||||
}
|
||||
|
||||
func TestDiscoverProviderModelsPreservesCorrelatedProxyError(t *testing.T) {
|
||||
manager, mockStore, mockPermissions, mockProxy := newModelDiscoveryManager(t)
|
||||
allowModelDiscovery(mockPermissions)
|
||||
mockStore.EXPECT().
|
||||
GetAgentNetworkProviderByID(gomock.Any(), store.LockingStrengthNone, "account-1", "provider-1").
|
||||
Return(&types.Provider{
|
||||
ID: "provider-1",
|
||||
AccountID: "account-1",
|
||||
ProviderID: "ollama",
|
||||
UpstreamURL: "http://ollama.internal:11434",
|
||||
}, nil)
|
||||
mockStore.EXPECT().
|
||||
GetAgentNetworkSettings(gomock.Any(), store.LockingStrengthNone, "account-1").
|
||||
Return(&types.Settings{AccountID: "account-1", Cluster: "private.example.com"}, nil)
|
||||
mockProxy.EXPECT().
|
||||
DiscoverModels(gomock.Any(), "account-1", "private.example.com", gomock.Any()).
|
||||
Return(&proto.ModelDiscoveryResult{
|
||||
RequestId: "probe-failed",
|
||||
Error: "upstream returned HTTP 401",
|
||||
}, nil)
|
||||
|
||||
_, err := manager.DiscoverProviderModels(context.Background(), "account-1", "user-1", "provider-1")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "upstream returned HTTP 401")
|
||||
assert.Contains(t, err.Error(), "probe-failed")
|
||||
}
|
||||
|
||||
func TestModelDiscoveryAdmissionControl(t *testing.T) {
|
||||
manager := &managerImpl{}
|
||||
start := time.Now()
|
||||
require.True(t, manager.beginModelDiscovery("account/provider", start))
|
||||
assert.False(t, manager.beginModelDiscovery("account/provider", start.Add(time.Second)))
|
||||
|
||||
manager.finishModelDiscovery("account/provider")
|
||||
assert.False(t, manager.beginModelDiscovery("account/provider", start.Add(time.Second)))
|
||||
assert.True(t, manager.beginModelDiscovery("account/provider", start.Add(modelDiscoveryCooldown)))
|
||||
}
|
||||
|
||||
func TestModelDiscoveryAdmissionControlExpiresFinishedAttempts(t *testing.T) {
|
||||
manager := &managerImpl{}
|
||||
start := time.Now()
|
||||
require.True(t, manager.beginModelDiscovery("account/provider", start))
|
||||
manager.finishModelDiscovery("account/provider")
|
||||
|
||||
manager.expireModelDiscoveryAttempt("account/provider", start)
|
||||
|
||||
manager.discoveryMu.Lock()
|
||||
_, retained := manager.discoveryAttempts["account/provider"]
|
||||
manager.discoveryMu.Unlock()
|
||||
assert.False(t, retained)
|
||||
}
|
||||
|
||||
func TestModelDiscoveryAdmissionControlKeepsReplacementAttempt(t *testing.T) {
|
||||
manager := &managerImpl{}
|
||||
start := time.Now()
|
||||
require.True(t, manager.beginModelDiscovery("account/provider", start))
|
||||
manager.finishModelDiscovery("account/provider")
|
||||
require.True(t, manager.beginModelDiscovery("account/provider", start.Add(modelDiscoveryCooldown)))
|
||||
|
||||
manager.expireModelDiscoveryAttempt("account/provider", start)
|
||||
|
||||
manager.discoveryMu.Lock()
|
||||
attempt, retained := manager.discoveryAttempts["account/provider"]
|
||||
manager.discoveryMu.Unlock()
|
||||
require.True(t, retained)
|
||||
assert.True(t, attempt.inFlight)
|
||||
assert.Equal(t, start.Add(modelDiscoveryCooldown), attempt.lastStarted)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package types
|
||||
|
||||
// DiscoveredModel is a normalized model identifier returned by a provider's
|
||||
// persisted upstream endpoint. Discovery does not persist models; the caller
|
||||
// must explicitly update the provider to opt into the returned list.
|
||||
type DiscoveredModel struct {
|
||||
ID string
|
||||
Label string
|
||||
}
|
||||
|
||||
// ModelDiscoveryResult describes a successful proxy-executed discovery probe.
|
||||
// RequestID correlates the management request with the proxy control message,
|
||||
// and ProxyCluster identifies the network vantage point that ran it.
|
||||
type ModelDiscoveryResult struct {
|
||||
RequestID string
|
||||
Source string
|
||||
ProxyCluster string
|
||||
Models []DiscoveredModel
|
||||
}
|
||||
Reference in New Issue
Block a user