mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-02 03:29:07 +02:00
Merge branch 'main' into reverse-proxy-crowdsec-appsec
This commit is contained in:
@@ -113,8 +113,61 @@ type Provider struct {
|
||||
// upstream provider + credentials on Portkey's hosted side).
|
||||
ExtraHeaders []ExtraHeader
|
||||
Models []Model
|
||||
// Discovery, when non-nil, describes how to ask this vendor which
|
||||
// models the operator's own credential can actually reach, so the
|
||||
// provider form can offer a live list instead of only the hand-curated
|
||||
// Models above. Nil for entries with no listing endpoint (gateways
|
||||
// vary too much) — those keep free-text entry.
|
||||
Discovery *Discovery
|
||||
}
|
||||
|
||||
// ListingShape names the response envelope a vendor returns its model
|
||||
// listing in. Every vendor invented its own, and none of them can be
|
||||
// guessed from the request, so the catalog states it.
|
||||
type ListingShape string
|
||||
|
||||
const (
|
||||
// ShapeOpenAIData is {"data":[{"id":…}]} — OpenAI, and Anthropic, which
|
||||
// adopted the same envelope.
|
||||
ShapeOpenAIData ListingShape = "openai_data"
|
||||
// ShapeBedrockInferenceProfiles is
|
||||
// {"inferenceProfileSummaries":[{"inferenceProfileId":…}]}. The ids carry
|
||||
// the region prefix that makes them invocable, which is exactly what an
|
||||
// operator cannot reconstruct by hand.
|
||||
ShapeBedrockInferenceProfiles ListingShape = "bedrock_inference_profiles"
|
||||
// ShapeVertexPublisherModels is {"publisherModels":[{"name":…}]}, where
|
||||
// name is a resource path and the invocable id is its last segment joined
|
||||
// to a separate versionId field.
|
||||
ShapeVertexPublisherModels ListingShape = "vertex_publisher_models"
|
||||
)
|
||||
|
||||
// Discovery describes one vendor's model-listing endpoint.
|
||||
//
|
||||
// Host is deliberately separate from the provider record's upstream URL:
|
||||
// Bedrock serves listings from the control plane (bedrock.<region>) while
|
||||
// inference must go to the runtime host (bedrock-runtime.<region>), so the
|
||||
// two cannot be the same value. Empty Host means "use the record's own
|
||||
// upstream", which is right for every vendor that serves both from one host.
|
||||
//
|
||||
// The regionPlaceholder in Host is substituted from the provider record's
|
||||
// region. Deriving the discovery host from the catalog rather than accepting
|
||||
// one from the caller is also what keeps this from being an open proxy: the
|
||||
// only hosts management will dial are the ones written here.
|
||||
type Discovery struct {
|
||||
Host string
|
||||
Path string
|
||||
Query string
|
||||
Shape ListingShape
|
||||
// Headers are static headers the vendor requires beyond the credential
|
||||
// (Anthropic versions its API through one and rejects a request without
|
||||
// it). The auth header itself comes from AuthHeaderName/Template.
|
||||
Headers map[string]string
|
||||
}
|
||||
|
||||
// RegionPlaceholder is replaced in Discovery.Host by the provider record's
|
||||
// configured region.
|
||||
const RegionPlaceholder = "<region>"
|
||||
|
||||
// ExtraHeader names a single optional per-provider routing/config
|
||||
// header. Catalog declares N of these per provider type; the operator
|
||||
// fills any subset on the provider record (see Provider.ExtraValues).
|
||||
@@ -245,8 +298,12 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#10A37F",
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
Discovery: &Discovery{
|
||||
Path: "/v1/models",
|
||||
Shape: ShapeOpenAIData,
|
||||
},
|
||||
ParserID: "openai",
|
||||
PricingSurfaces: []string{"openai"},
|
||||
// Pricing + context windows cross-checked against LiteLLM's
|
||||
// model_prices_and_context_window.json. Notable corrections from
|
||||
// earlier values: o4-mini repriced from $4/$16 to $1.10/$4.40
|
||||
@@ -284,8 +341,18 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#D97757",
|
||||
ParserID: "anthropic",
|
||||
PricingSurfaces: []string{"anthropic"},
|
||||
Discovery: &Discovery{
|
||||
Path: "/v1/models",
|
||||
// The default page is short and a picker wants the whole
|
||||
// catalogue in one call.
|
||||
Query: "limit=1000",
|
||||
Shape: ShapeOpenAIData,
|
||||
// Anthropic versions its API through a header and refuses a
|
||||
// request that omits it, listing included.
|
||||
Headers: map[string]string{"anthropic-version": "2023-06-01"},
|
||||
},
|
||||
ParserID: "anthropic",
|
||||
PricingSurfaces: []string{"anthropic"},
|
||||
// Per Anthropic's current model lineup. Pricing in USD per 1k
|
||||
// tokens. Context windows: 4.6+ family is 1M; Haiku 4.5 stays at
|
||||
// 200K. claude-3-7-sonnet and claude-3-5-haiku retired
|
||||
@@ -296,6 +363,8 @@ var providers = []Provider{
|
||||
// account to be on >= 30-day data retention or all requests
|
||||
// 400.
|
||||
Models: []Model{
|
||||
{ID: "claude-opus-5", Label: "Claude Opus 5", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-5", Label: "Claude Sonnet 5", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
@@ -343,6 +412,22 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#FF9900",
|
||||
// Listings come from the CONTROL PLANE, not the runtime host in
|
||||
// DefaultHost above: ListInferenceProfiles is not an operation
|
||||
// bedrock-runtime implements, and answers <UnknownOperationException/>
|
||||
// there. Inference has to go to the runtime host, so the two hosts
|
||||
// genuinely differ and Discovery.Host carries the difference.
|
||||
//
|
||||
// Inference profiles rather than foundation models because the profile
|
||||
// id is the invocable one: it carries the region prefix (eu., us.,
|
||||
// global.) that AWS requires and that cannot be derived from the
|
||||
// configured region — an eu-central-1 account legitimately holds
|
||||
// global.* profiles.
|
||||
Discovery: &Discovery{
|
||||
Host: "bedrock." + RegionPlaceholder + ".amazonaws.com",
|
||||
Path: "/inference-profiles",
|
||||
Shape: ShapeBedrockInferenceProfiles,
|
||||
},
|
||||
// ParserID stays empty (path-style dispatch via IsBedrockPathStyle);
|
||||
// the request parser meters these under the "bedrock" surface.
|
||||
PricingSurfaces: []string{"bedrock"},
|
||||
@@ -355,6 +440,8 @@ var providers = []Provider{
|
||||
// Llama 3.3 70B entry kept unchanged — LiteLLM tracks only
|
||||
// per-region Llama 3 entries; standalone 3.3 not yet listed.
|
||||
Models: []Model{
|
||||
{ID: "anthropic.claude-opus-5", Label: "Claude Opus 5 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-sonnet-5", Label: "Claude Sonnet 5 (Bedrock)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-8", Label: "Claude Opus 4.8 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-7", Label: "Claude Opus 4.7 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "anthropic.claude-opus-4-6", Label: "Claude Opus 4.6 (Bedrock)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
@@ -391,6 +478,15 @@ var providers = []Provider{
|
||||
AuthHeaderTemplate: "Bearer ${API_KEY}",
|
||||
DefaultContentType: "application/json",
|
||||
BrandColor: "#4285F4",
|
||||
// Only the v1beta1 publisher listing answers: the v1 form and the
|
||||
// project-scoped form under BOTH versions return 404. That means the
|
||||
// list is publisher-global — it cannot say which models this project
|
||||
// has enabled — so it is offered as a suggestion beside the catalog
|
||||
// rather than replacing it. See the discovery e2e for the probes.
|
||||
Discovery: &Discovery{
|
||||
Path: "/v1beta1/publishers/anthropic/models",
|
||||
Shape: ShapeVertexPublisherModels,
|
||||
},
|
||||
// ParserID stays empty (path-style dispatch via IsVertexPathStyle);
|
||||
// Anthropic-on-Vertex requests are metered under the "anthropic"
|
||||
// surface with the bare, unversioned model id.
|
||||
@@ -406,6 +502,8 @@ var providers = []Provider{
|
||||
// exists — the router denies unmeterable publishers rather than forward
|
||||
// them uncounted.
|
||||
Models: []Model{
|
||||
{ID: "claude-opus-5", Label: "Claude Opus 5 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-sonnet-5", Label: "Claude Sonnet 5 (Vertex)", InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003, CacheCreationPer1k: 0.00375, ContextWindow: 1000000},
|
||||
{ID: "claude-fable-5", Label: "Claude Fable 5 (Vertex)", InputPer1k: 0.010, OutputPer1k: 0.050, CacheReadPer1k: 0.001, CacheCreationPer1k: 0.0125, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-8", Label: "Claude Opus 4.8 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
{ID: "claude-opus-4-7", Label: "Claude Opus 4.7 (Vertex)", InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625, ContextWindow: 1000000},
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
package catalog
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// TestClaudeLineupSelectable pins the models Claude Code resolves to by
|
||||
// default. A model absent from the lineup can't be ticked on a provider
|
||||
// record, so llm_router denies it as not-routable and the operator has no
|
||||
// way to authorise the client's own default.
|
||||
func TestClaudeLineupSelectable(t *testing.T) {
|
||||
for providerID, wanted := range map[string][]string{
|
||||
"anthropic_api": {"claude-opus-5", "claude-sonnet-5", "claude-haiku-4-5"},
|
||||
"bedrock_api": {"anthropic.claude-opus-5", "anthropic.claude-sonnet-5", "anthropic.claude-haiku-4-5"},
|
||||
"vertex_ai_api": {"claude-opus-5", "claude-sonnet-5", "claude-haiku-4-5"},
|
||||
} {
|
||||
provider, ok := Lookup(providerID)
|
||||
require.True(t, ok, "catalog must define %s", providerID)
|
||||
|
||||
selectable := make(map[string]Model, len(provider.Models))
|
||||
for _, m := range provider.Models {
|
||||
selectable[m.ID] = m
|
||||
}
|
||||
for _, id := range wanted {
|
||||
model, found := selectable[id]
|
||||
require.True(t, found, "%s must offer %s", providerID, id)
|
||||
assert.NotEmpty(t, model.Label, "%s/%s needs a label for the picker", providerID, id)
|
||||
assert.Positive(t, model.InputPer1k, "%s/%s needs an input rate", providerID, id)
|
||||
assert.Positive(t, model.OutputPer1k, "%s/%s needs an output rate", providerID, id)
|
||||
assert.Positive(t, model.ContextWindow, "%s/%s needs a context window", providerID, id)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -0,0 +1,178 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/shared/auth"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// discoveryManagerStub records what the handler asked for and returns a canned
|
||||
// answer. The Manager interface is embedded rather than implemented: only the
|
||||
// one method is reachable from this handler, and a call to any other should
|
||||
// fail loudly rather than silently return a zero value.
|
||||
type discoveryManagerStub struct {
|
||||
agentnetwork.Manager
|
||||
|
||||
gotReq modeldiscovery.Request
|
||||
gotRecordID string
|
||||
models []modeldiscovery.Model
|
||||
err error
|
||||
}
|
||||
|
||||
func (s *discoveryManagerStub) DiscoverProviderModels(
|
||||
_ context.Context, _, _ string, req modeldiscovery.Request, recordID string,
|
||||
) ([]modeldiscovery.Model, error) {
|
||||
s.gotReq = req
|
||||
s.gotRecordID = recordID
|
||||
return s.models, s.err
|
||||
}
|
||||
|
||||
// postDiscovery drives the handler with an authenticated request.
|
||||
func postDiscovery(t *testing.T, stub *discoveryManagerStub, body string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
h := &handler{manager: stub}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/agent-network/catalog/providers/models", strings.NewReader(body))
|
||||
req = req.WithContext(nbcontext.SetUserAuthInContext(req.Context(), auth.UserAuth{
|
||||
AccountId: "acc-1",
|
||||
UserId: "user-1",
|
||||
}))
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
h.discoverProviderModels(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
func TestDiscoverModelsReturnsTheVendorList(t *testing.T) {
|
||||
stub := &discoveryManagerStub{models: []modeldiscovery.Model{
|
||||
{ID: "eu.anthropic.claude-haiku-4-5-20251001-v1:0", Label: "EU Claude Haiku 4.5", PricingKnown: true},
|
||||
{ID: "global.cohere.embed-v4:0", Label: "Global Cohere Embed v4"},
|
||||
// A vendor that supplies no display name at all. Bedrock does for
|
||||
// every profile, but the OpenAI listing carries none.
|
||||
{ID: "gpt-4o-mini", PricingKnown: true},
|
||||
}}
|
||||
|
||||
rec := postDiscovery(t, stub, `{
|
||||
"catalog_provider_id":"bedrock_api",
|
||||
"upstream_url":"https://bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
"api_key":"aws-bearer"
|
||||
}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String())
|
||||
|
||||
var out api.AgentNetworkModelDiscoveryResponse
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &out))
|
||||
require.Len(t, out.Models, 3)
|
||||
|
||||
assert.Equal(t, "eu.anthropic.claude-haiku-4-5-20251001-v1:0", out.Models[0].Id)
|
||||
assert.True(t, out.Models[0].PricingKnown)
|
||||
// An unpriced model must say so rather than arriving indistinguishable
|
||||
// from a priced one: registering it silently would meter at zero.
|
||||
assert.False(t, out.Models[1].PricingKnown)
|
||||
|
||||
require.NotNil(t, out.Models[0].Label, "the vendor supplied a display name")
|
||||
assert.Equal(t, "EU Claude Haiku 4.5", *out.Models[0].Label)
|
||||
// A vendor that supplies no name must omit the key rather than send an
|
||||
// empty string: the dashboard falls back to the id on absence, and would
|
||||
// render a blank row for "".
|
||||
assert.Nil(t, out.Models[2].Label, "an absent label must not serialize")
|
||||
assert.NotContains(t, rec.Body.String(), `"label":""`)
|
||||
|
||||
assert.Equal(t, "bedrock_api", stub.gotReq.CatalogID)
|
||||
assert.Equal(t, "aws-bearer", stub.gotReq.APIKey)
|
||||
// The upstream is what the region is read back out of for Bedrock, so
|
||||
// losing it here would break discovery for every regional provider.
|
||||
assert.Equal(t, "https://bedrock-runtime.eu-central-1.amazonaws.com", stub.gotReq.UpstreamURL)
|
||||
assert.Empty(t, stub.gotRecordID)
|
||||
}
|
||||
|
||||
func TestDiscoverModelsUsesAStoredRecordWithoutAKey(t *testing.T) {
|
||||
stub := &discoveryManagerStub{}
|
||||
|
||||
rec := postDiscovery(t, stub, `{"catalog_provider_id":"openai_api","provider_id":"prov-42"}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String())
|
||||
|
||||
// The dashboard refreshes a saved provider's list without ever holding
|
||||
// the credential, so the record id has to reach the manager.
|
||||
assert.Equal(t, "prov-42", stub.gotRecordID)
|
||||
assert.Empty(t, stub.gotReq.APIKey)
|
||||
}
|
||||
|
||||
// TestDiscoverModelsRefusesMixedCredentials covers the case where a caller
|
||||
// names a saved provider AND supplies a key. Accepting it would run an
|
||||
// arbitrary credential under the identity of a record the caller may only be
|
||||
// permitted to read.
|
||||
func TestDiscoverModelsRefusesMixedCredentials(t *testing.T) {
|
||||
stub := &discoveryManagerStub{}
|
||||
|
||||
rec := postDiscovery(t, stub, `{
|
||||
"catalog_provider_id":"openai_api",
|
||||
"provider_id":"prov-42",
|
||||
"api_key":"sk-attacker"
|
||||
}`)
|
||||
assert.Equal(t, http.StatusBadRequest, rec.Code)
|
||||
assert.Empty(t, stub.gotRecordID, "the request must be refused before it reaches the manager")
|
||||
}
|
||||
|
||||
// TestDiscoverModelsReportsNoDiscoveryDistinctly matters because the caller
|
||||
// falls back to the catalog's own model list on this outcome. Collapsing it
|
||||
// into a generic 500 would turn "this provider has no listing endpoint" into
|
||||
// "something went wrong", and the form would show an error instead of a list.
|
||||
func TestDiscoverModelsReportsNoDiscoveryDistinctly(t *testing.T) {
|
||||
stub := &discoveryManagerStub{err: modeldiscovery.ErrNoDiscovery}
|
||||
|
||||
rec := postDiscovery(t, stub, `{"catalog_provider_id":"litellm_proxy","upstream_url":"https://gw.example.com","api_key":"sk"}`)
|
||||
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code)
|
||||
}
|
||||
|
||||
// TestDiscoverModelsTrimsTheCatalogID pins that the id the emptiness check
|
||||
// accepts is the id the manager receives. A padded value that clears the check
|
||||
// but reaches the catalog untrimmed misses the lookup, and the operator is told
|
||||
// their provider does not exist.
|
||||
func TestDiscoverModelsTrimsTheCatalogID(t *testing.T) {
|
||||
stub := &discoveryManagerStub{}
|
||||
|
||||
rec := postDiscovery(t, stub, `{"catalog_provider_id":" openai_api ","api_key":"sk"}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "body: %s", rec.Body.String())
|
||||
assert.Equal(t, "openai_api", stub.gotReq.CatalogID)
|
||||
}
|
||||
|
||||
// TestDiscoverModelsReportsCallerInputAsBadRequest covers the other half of the
|
||||
// error mapping. These failures are all reachable from a well-formed request
|
||||
// with a bad field value, so answering 500 both misinforms the operator and
|
||||
// puts their typo into the server's error rate.
|
||||
func TestDiscoverModelsReportsCallerInputAsBadRequest(t *testing.T) {
|
||||
stub := &discoveryManagerStub{
|
||||
err: fmt.Errorf("%w: unknown catalog provider %q", modeldiscovery.ErrInvalidRequest, "nope"),
|
||||
}
|
||||
|
||||
rec := postDiscovery(t, stub, `{"catalog_provider_id":"nope","api_key":"sk"}`)
|
||||
assert.Equal(t, http.StatusBadRequest, rec.Code)
|
||||
assert.Contains(t, rec.Body.String(), "unknown catalog provider")
|
||||
}
|
||||
|
||||
func TestDiscoverModelsRejectsMalformedRequests(t *testing.T) {
|
||||
for name, body := range map[string]string{
|
||||
"not json": `{`,
|
||||
"no catalog provider": `{"api_key":"sk"}`,
|
||||
"blank catalog provider": `{"catalog_provider_id":" ","api_key":"sk"}`,
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
stub := &discoveryManagerStub{}
|
||||
rec := postDiscovery(t, stub, body)
|
||||
assert.Equal(t, http.StatusBadRequest, rec.Code)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,7 @@ package handlers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/url"
|
||||
@@ -16,6 +17,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
@@ -32,6 +34,7 @@ type handler struct {
|
||||
func RegisterEndpoints(manager agentnetwork.Manager, router *mux.Router) {
|
||||
h := &handler{manager: manager}
|
||||
router.HandleFunc("/agent-network/catalog/providers", h.getCatalogProviders).Methods("GET", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/catalog/providers/models", h.discoverProviderModels).Methods("POST", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/providers", h.getAllProviders).Methods("GET", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET", "OPTIONS")
|
||||
@@ -61,6 +64,98 @@ func (h *handler) getCatalogProviders(w http.ResponseWriter, r *http.Request) {
|
||||
util.WriteJSONObject(r.Context(), w, out)
|
||||
}
|
||||
|
||||
// discoverProviderModels asks the vendor which models the operator's own
|
||||
// credential can reach, so the provider form can offer a live list rather than
|
||||
// only the static catalog.
|
||||
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
|
||||
}
|
||||
|
||||
var body api.AgentNetworkModelDiscoveryRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||
util.WriteErrorResponse("invalid json", http.StatusBadRequest, w)
|
||||
return
|
||||
}
|
||||
// Trimmed once and carried, not trimmed for the emptiness test and then
|
||||
// discarded: a padded " openai_api " would clear the check here and miss
|
||||
// the catalog lookup, reporting the provider as unknown.
|
||||
catalogID := strings.TrimSpace(body.CatalogProviderId)
|
||||
if catalogID == "" {
|
||||
util.WriteErrorResponse("catalog_provider_id is required", http.StatusBadRequest, w)
|
||||
return
|
||||
}
|
||||
|
||||
recordID := strValue(body.ProviderId)
|
||||
req := modeldiscovery.Request{
|
||||
CatalogID: catalogID,
|
||||
UpstreamURL: strValue(body.UpstreamUrl),
|
||||
APIKey: strValue(body.ApiKey),
|
||||
}
|
||||
// One source of credential or the other, never a mix: taking a key from
|
||||
// the request while addressing a saved record would let a caller run an
|
||||
// arbitrary credential against a provider they can only read.
|
||||
if recordID != "" && req.APIKey != "" {
|
||||
util.WriteErrorResponse("provide either provider_id or api_key, not both", http.StatusBadRequest, w)
|
||||
return
|
||||
}
|
||||
|
||||
models, err := h.manager.DiscoverProviderModels(r.Context(), userAuth.AccountId, userAuth.UserId, req, recordID)
|
||||
if err != nil {
|
||||
// A provider with no listing endpoint is a fact about the catalog
|
||||
// entry, not a failure: the caller falls back to the catalog's own
|
||||
// models, so it must be able to tell the two apart.
|
||||
if errors.Is(err, modeldiscovery.ErrNoDiscovery) {
|
||||
util.WriteErrorResponse(err.Error(), http.StatusUnprocessableEntity, w)
|
||||
return
|
||||
}
|
||||
// An unknown provider, an unusable upstream, a missing region or a
|
||||
// missing key are all things the caller sent, reachable from a
|
||||
// well-formed request. Reporting them as 500 tells the operator the
|
||||
// server broke and buries genuine faults in the error rate.
|
||||
if errors.Is(err, modeldiscovery.ErrInvalidRequest) {
|
||||
util.WriteErrorResponse(err.Error(), http.StatusBadRequest, w)
|
||||
return
|
||||
}
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
out := api.AgentNetworkModelDiscoveryResponse{Models: make([]api.AgentNetworkDiscoveredModel, 0, len(models))}
|
||||
for _, m := range models {
|
||||
entry := api.AgentNetworkDiscoveredModel{
|
||||
Id: m.ID,
|
||||
PricingKnown: m.PricingKnown,
|
||||
// Sent even when zero: the form prefills every discovered model as
|
||||
// an editable row, and an unpriced one is shown at zero and flagged
|
||||
// rather than left out.
|
||||
InputPer1k: m.InputPer1k,
|
||||
OutputPer1k: m.OutputPer1k,
|
||||
// Cache rates stay absent when unset, matching the catalog
|
||||
// response — a zero would read as "free", not "not applicable".
|
||||
CachedInputPer1k: positiveRatePtr(m.CachedInputPer1k),
|
||||
CacheReadPer1k: positiveRatePtr(m.CacheReadPer1k),
|
||||
CacheCreationPer1k: positiveRatePtr(m.CacheCreationPer1k),
|
||||
}
|
||||
if m.Label != "" {
|
||||
label := m.Label
|
||||
entry.Label = &label
|
||||
}
|
||||
out.Models = append(out.Models, entry)
|
||||
}
|
||||
util.WriteJSONObject(r.Context(), w, out)
|
||||
}
|
||||
|
||||
// strValue reads an optional string field, treating absent as empty.
|
||||
func strValue(v *string) string {
|
||||
if v == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(*v)
|
||||
}
|
||||
|
||||
// applyDefaultPricing overwrites the catalog response's model rates with
|
||||
// the LIVE default pricing table, which may differ from the compiled-in
|
||||
// catalog rates when the operator provides a defaults_llm_pricing.yaml.
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/labelgen"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
|
||||
@@ -50,6 +51,7 @@ type Manager interface {
|
||||
CreateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
|
||||
UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
|
||||
DeleteProvider(ctx context.Context, accountID, userID, providerID string) error
|
||||
DiscoverProviderModels(ctx context.Context, accountID, userID string, req modeldiscovery.Request, recordID string) ([]modeldiscovery.Model, error)
|
||||
|
||||
GetAllPolicies(ctx context.Context, accountID, userID string) ([]*types.Policy, error)
|
||||
GetPolicy(ctx context.Context, accountID, userID, policyID string) (*types.Policy, error)
|
||||
@@ -123,6 +125,15 @@ type managerImpl struct {
|
||||
permissionsManager permissions.Manager
|
||||
proxyController proxy.Controller
|
||||
|
||||
// modelDiscovery queries vendors for the models a credential can reach.
|
||||
// A field rather than a package call so tests can drive it without
|
||||
// reaching the network.
|
||||
//
|
||||
// One instance serves every request for the process's lifetime, so its
|
||||
// fields must stay read-only after construction: lazy initialisation
|
||||
// inside Fetch or httpClient would race across request goroutines.
|
||||
modelDiscovery *modeldiscovery.Client
|
||||
|
||||
// reconcileCache holds the last set of synthesised proxy mappings
|
||||
// per account, each paired with the proxy that served it, so a change
|
||||
// of serving proxy can be diffed without re-deriving it.
|
||||
@@ -151,6 +162,7 @@ func NewManager(
|
||||
accountManager: accountManager,
|
||||
permissionsManager: permissionsManager,
|
||||
proxyController: proxyController,
|
||||
modelDiscovery: &modeldiscovery.Client{},
|
||||
reconcileCache: make(map[string]map[string]syntheticMapping),
|
||||
labelRng: rand.New(rand.NewSource(time.Now().UnixNano())),
|
||||
}
|
||||
@@ -170,6 +182,38 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid
|
||||
return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
|
||||
}
|
||||
|
||||
// DiscoverProviderModels asks the vendor which models a credential can reach.
|
||||
//
|
||||
// recordID, when set, names an existing provider whose stored credential and
|
||||
// upstream are used instead of the ones in req — so the dashboard can refresh
|
||||
// the list without ever holding the key.
|
||||
//
|
||||
// Gated on Create rather than Read: this spends the operator's credential
|
||||
// against a third party, which is not something a read-only role should be
|
||||
// able to make the server do. That one check also covers reading the stored
|
||||
// record — Create is strictly stronger than Read here, and the lookup is
|
||||
// scoped to accountID, so another account's record is never reachable.
|
||||
func (m *managerImpl) DiscoverProviderModels(ctx context.Context, accountID, userID string, req modeldiscovery.Request, recordID string) ([]modeldiscovery.Model, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Create); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if recordID != "" {
|
||||
record, err := m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, recordID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// The catalog id comes from the stored record too: letting the caller
|
||||
// name a different one would run a provider's credential against
|
||||
// whichever vendor endpoint they picked.
|
||||
req.CatalogID = record.ProviderID
|
||||
req.UpstreamURL = record.UpstreamURL
|
||||
req.APIKey = record.APIKey
|
||||
}
|
||||
|
||||
return m.modelDiscovery.Fetch(ctx, req)
|
||||
}
|
||||
|
||||
// CreateProvider persists a new provider for the account. Providers have no
|
||||
// settings side effects: the account's endpoint is bootstrapped separately and
|
||||
// explicitly via CreateSettings, and every provider in the account routes
|
||||
@@ -1017,6 +1061,10 @@ func (*mockManager) GetAllProviders(_ context.Context, _, _ string) ([]*types.Pr
|
||||
return []*types.Provider{}, nil
|
||||
}
|
||||
|
||||
func (*mockManager) DiscoverProviderModels(_ context.Context, _, _ string, _ modeldiscovery.Request, _ string) ([]modeldiscovery.Model, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (*mockManager) GetProvider(_ context.Context, _, _, _ string) (*types.Provider, error) {
|
||||
return &types.Provider{}, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,469 @@
|
||||
// Package modeldiscovery asks a vendor which models an operator's own
|
||||
// credential can reach, so the provider form can offer a live list instead of
|
||||
// only the catalog's hand-curated one.
|
||||
//
|
||||
// The catalog cannot know two things that matter. It goes stale — its entries
|
||||
// carry comments tracking which models a vendor retired on which date — and it
|
||||
// cannot see an account: which OpenAI models an org is entitled to, which
|
||||
// Bedrock inference profiles a given account and region hold, which Vertex
|
||||
// models a project has enabled. Those are exactly the facts an operator needs
|
||||
// when filling in a provider record, and only the vendor has them.
|
||||
//
|
||||
// The vendor is authoritative for the model ID. The catalog remains
|
||||
// authoritative for pricing, and a discovered model the catalog cannot price
|
||||
// is reported as such rather than silently registered at a rate of zero.
|
||||
package modeldiscovery
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"golang.org/x/oauth2/google"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
)
|
||||
|
||||
const (
|
||||
// fetchTimeout bounds one vendor call end to end. A listing is a single
|
||||
// small GET; anything slower is a vendor problem and the operator is
|
||||
// waiting on a form.
|
||||
fetchTimeout = 8 * time.Second
|
||||
// maxListingBytes bounds the response we will buffer. The largest real
|
||||
// listing observed is Bedrock's foundation-model catalogue at ~70KB, so
|
||||
// this is a wide margin over anything legitimate.
|
||||
maxListingBytes = 2 << 20
|
||||
// gcpScope matches the scope llm_router mints Vertex tokens under, so a
|
||||
// credential that works for discovery works for inference too.
|
||||
gcpScope = "https://www.googleapis.com/auth/cloud-platform"
|
||||
// vertexKeyfilePrefix marks an api_key that is a base64 service-account
|
||||
// JSON key rather than a bearer token.
|
||||
vertexKeyfilePrefix = "keyfile::"
|
||||
)
|
||||
|
||||
// ErrNoDiscovery is returned for a catalog entry that declares no listing
|
||||
// endpoint. Gateways vary too much to have one, and the caller should fall
|
||||
// back to the catalog list plus free-text entry rather than treating this as
|
||||
// a failure.
|
||||
var ErrNoDiscovery = errors.New("provider has no model-discovery endpoint")
|
||||
|
||||
// ErrInvalidRequest marks a discovery failure caused by the caller's own input
|
||||
// rather than by the vendor or by this server. Every one of these is reachable
|
||||
// from a well-formed request carrying a bad field value, so the handler owes
|
||||
// the caller a 400 — a 500 would both misinform them and bury real server
|
||||
// faults in the error rate.
|
||||
var ErrInvalidRequest = errors.New("invalid discovery request")
|
||||
|
||||
// Model is one discovered model.
|
||||
type Model struct {
|
||||
// ID is the identifier to register on the provider record, in the form the
|
||||
// vendor issues it. For Bedrock that is the region-prefixed inference
|
||||
// profile id, which is the only form AWS accepts at invoke time.
|
||||
ID string
|
||||
// Label is the vendor's display name where it supplies one.
|
||||
Label string
|
||||
// PricingKnown reports whether the shipped pricing table can price this
|
||||
// model. False means the operator must set rates, or the request would
|
||||
// meter at zero.
|
||||
PricingKnown bool
|
||||
// The rates below are the defaults for this model, taken from the same
|
||||
// table the proxy bills with, so the form prefills exactly what a request
|
||||
// would cost. All zero when PricingKnown is false — an unpriced model is
|
||||
// offered at zero and flagged, rather than withheld: the vendor says the
|
||||
// credential can reach it, and refusing to show it would hide a model the
|
||||
// operator genuinely has.
|
||||
InputPer1k float64
|
||||
OutputPer1k float64
|
||||
CachedInputPer1k float64
|
||||
CacheReadPer1k float64
|
||||
CacheCreationPer1k float64
|
||||
}
|
||||
|
||||
// Request identifies which vendor to ask and with what credential.
|
||||
type Request struct {
|
||||
// CatalogID selects the catalog entry, which supplies the endpoint, the
|
||||
// auth header and the response shape. The caller never supplies those.
|
||||
CatalogID string
|
||||
// UpstreamURL is the provider record's configured upstream. It is used
|
||||
// only when the catalog entry declares no discovery host of its own.
|
||||
UpstreamURL string
|
||||
// Region substitutes the catalog host's <region> placeholder.
|
||||
Region string
|
||||
// APIKey is the operator's credential, exactly as stored on the record.
|
||||
APIKey string
|
||||
}
|
||||
|
||||
// Client fetches model listings. The zero value is usable; Resolver and
|
||||
// HTTPClient exist so tests can drive it against a local server.
|
||||
type Client struct {
|
||||
HTTPClient *http.Client
|
||||
// Resolver looks up the host for the SSRF check. Nil uses the default.
|
||||
Resolver *net.Resolver
|
||||
// AllowPrivateHosts disables the private-address guard. Only tests set it:
|
||||
// their server is on loopback, which is precisely what the guard blocks.
|
||||
AllowPrivateHosts bool
|
||||
}
|
||||
|
||||
// Fetch returns the models the credential can reach.
|
||||
func (c *Client) Fetch(ctx context.Context, req Request) ([]Model, error) {
|
||||
entry, ok := catalog.Lookup(req.CatalogID)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("%w: unknown catalog provider %q", ErrInvalidRequest, req.CatalogID)
|
||||
}
|
||||
if entry.Discovery == nil {
|
||||
return nil, ErrNoDiscovery
|
||||
}
|
||||
|
||||
endpoint, err := c.discoveryURL(entry, req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, fetchTimeout)
|
||||
defer cancel()
|
||||
|
||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build discovery request: %w", err)
|
||||
}
|
||||
if err := applyAuth(httpReq, entry, req.APIKey); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for name, value := range entry.Discovery.Headers {
|
||||
httpReq.Header.Set(name, value)
|
||||
}
|
||||
httpReq.Header.Set("Accept", "application/json")
|
||||
|
||||
resp, err := c.httpClient().Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reach %s: %w", entry.Name, err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxListingBytes))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read %s listing: %w", entry.Name, err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
// Surface the vendor's own status. An operator whose key lacks a scope
|
||||
// needs to see 403 rather than a generic failure.
|
||||
return nil, fmt.Errorf("%s returned %d for its model listing", entry.Name, resp.StatusCode)
|
||||
}
|
||||
|
||||
ids, err := parseListing(entry.Discovery.Shape, body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return decorate(entry, ids), nil
|
||||
}
|
||||
|
||||
// discoveryURL builds the listing URL and refuses one that does not point at a
|
||||
// public host.
|
||||
//
|
||||
// The path, query and (for Bedrock) the host all come from the catalog rather
|
||||
// than from the caller, so the only operator-controlled part is the host of an
|
||||
// entry whose listing lives on its own upstream. That still has to be checked:
|
||||
// management holds credentials for every provider, and an upstream pointed at
|
||||
// an internal address would turn this endpoint into a probe of the management
|
||||
// server's own network.
|
||||
func (c *Client) discoveryURL(entry catalog.Provider, req Request) (string, error) {
|
||||
host := entry.Discovery.Host
|
||||
if host == "" {
|
||||
parsed, err := url.Parse(strings.TrimSpace(req.UpstreamURL))
|
||||
if err != nil || parsed.Host == "" {
|
||||
return "", fmt.Errorf("%w: provider upstream %q is not a usable URL", ErrInvalidRequest, req.UpstreamURL)
|
||||
}
|
||||
host = parsed.Host
|
||||
}
|
||||
if strings.Contains(host, catalog.RegionPlaceholder) {
|
||||
region := strings.TrimSpace(req.Region)
|
||||
if region == "" {
|
||||
// A provider record carries no region field: the region lives
|
||||
// inside the upstream host the operator already configured, so
|
||||
// read it back out rather than asking them for it twice.
|
||||
region = RegionFromUpstream(entry, req.UpstreamURL)
|
||||
}
|
||||
if region == "" {
|
||||
return "", fmt.Errorf("%w: %s discovery needs a region, and none could be read from the provider upstream",
|
||||
ErrInvalidRequest, entry.Name)
|
||||
}
|
||||
host = strings.ReplaceAll(host, catalog.RegionPlaceholder, region)
|
||||
}
|
||||
|
||||
target := &url.URL{Scheme: "https", Host: host, Path: entry.Discovery.Path, RawQuery: entry.Discovery.Query}
|
||||
if err := c.checkPublicHost(target.Hostname()); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return target.String(), nil
|
||||
}
|
||||
|
||||
// RegionFromUpstream recovers the region an operator embedded in the provider
|
||||
// upstream, by matching it against the catalog's own host template. Bedrock's
|
||||
// template is "bedrock-runtime.<region>.amazonaws.com" and Vertex's is
|
||||
// "<region>-aiplatform.googleapis.com", so the region is whatever sits between
|
||||
// the fixed halves. Returns empty when the upstream does not match the
|
||||
// template, which is the case for a custom or proxied endpoint.
|
||||
func RegionFromUpstream(entry catalog.Provider, upstreamURL string) string {
|
||||
prefix, suffix, found := strings.Cut(entry.DefaultHost, catalog.RegionPlaceholder)
|
||||
if !found {
|
||||
return ""
|
||||
}
|
||||
parsed, err := url.Parse(strings.TrimSpace(upstreamURL))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
host := parsed.Hostname()
|
||||
if host == "" {
|
||||
// A bare host with no scheme parses as a path, not a host.
|
||||
host = strings.TrimSpace(upstreamURL)
|
||||
}
|
||||
// The two halves must not overlap. "bedrock-runtime.amazonaws.com" carries
|
||||
// both of Bedrock's — it is the regionless endpoint — and satisfies both
|
||||
// checks above while leaving nothing between them, so slicing it would
|
||||
// panic on an inverted range rather than report "no region here".
|
||||
if !strings.HasPrefix(host, prefix) || !strings.HasSuffix(host, suffix) ||
|
||||
len(host) < len(prefix)+len(suffix) {
|
||||
return ""
|
||||
}
|
||||
region := host[len(prefix) : len(host)-len(suffix)]
|
||||
if region == "" || strings.Contains(region, ".") {
|
||||
return ""
|
||||
}
|
||||
return region
|
||||
}
|
||||
|
||||
// checkPublicHost refuses hosts that resolve to an address the management
|
||||
// server should never be asked to reach on an operator's behalf.
|
||||
func (c *Client) checkPublicHost(host string) error {
|
||||
if c.AllowPrivateHosts {
|
||||
return nil
|
||||
}
|
||||
if host == "" {
|
||||
return errors.New("discovery host is empty")
|
||||
}
|
||||
resolver := c.Resolver
|
||||
if resolver == nil {
|
||||
resolver = net.DefaultResolver
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), fetchTimeout)
|
||||
defer cancel()
|
||||
|
||||
addrs, err := resolver.LookupNetIP(ctx, "ip", host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve discovery host %q: %w", host, err)
|
||||
}
|
||||
// Every address must be public: a name that resolves to one public and one
|
||||
// loopback address is still a way to reach loopback.
|
||||
for _, addr := range addrs {
|
||||
if !isPublic(addr) {
|
||||
return fmt.Errorf("%w: discovery host %q resolves to a non-public address", ErrInvalidRequest, host)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// isPublic reports whether an address is one we are willing to dial.
|
||||
func isPublic(addr netip.Addr) bool {
|
||||
addr = addr.Unmap()
|
||||
switch {
|
||||
case !addr.IsValid(),
|
||||
addr.IsLoopback(),
|
||||
addr.IsPrivate(),
|
||||
addr.IsLinkLocalUnicast(),
|
||||
addr.IsLinkLocalMulticast(),
|
||||
addr.IsInterfaceLocalMulticast(),
|
||||
addr.IsMulticast(),
|
||||
addr.IsUnspecified():
|
||||
return false
|
||||
}
|
||||
// 100.64.0.0/10 (carrier NAT) is where NetBird's own overlay addresses
|
||||
// live, so it is emphatically not somewhere to send a provider credential.
|
||||
if addr.Is4() {
|
||||
b := addr.As4()
|
||||
if b[0] == 100 && b[1] >= 64 && b[1] <= 127 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// applyAuth sets the credential header the catalog entry declares. A Vertex
|
||||
// service-account key is exchanged for an OAuth token first, the same way the
|
||||
// proxy does at request time.
|
||||
func applyAuth(req *http.Request, entry catalog.Provider, apiKey string) error {
|
||||
key := strings.TrimSpace(apiKey)
|
||||
if key == "" {
|
||||
return fmt.Errorf("%w: %s discovery needs an API key", ErrInvalidRequest, entry.Name)
|
||||
}
|
||||
if rest, ok := strings.CutPrefix(key, vertexKeyfilePrefix); ok {
|
||||
token, err := mintGCPToken(req.Context(), rest)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
key = token
|
||||
}
|
||||
name := entry.AuthHeaderName
|
||||
if name == "" {
|
||||
name = "Authorization"
|
||||
}
|
||||
template := entry.AuthHeaderTemplate
|
||||
if template == "" {
|
||||
template = "${API_KEY}"
|
||||
}
|
||||
req.Header.Set(name, strings.ReplaceAll(template, "${API_KEY}", key))
|
||||
return nil
|
||||
}
|
||||
|
||||
// mintGCPToken exchanges a base64 service-account key for an access token.
|
||||
func mintGCPToken(ctx context.Context, saKeyB64 string) (string, error) {
|
||||
jsonKey, err := base64.StdEncoding.DecodeString(strings.TrimSpace(saKeyB64))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("decode service-account key: %w", err)
|
||||
}
|
||||
conf, err := google.JWTConfigFromJSON(jsonKey, gcpScope)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("parse service-account key: %w", err)
|
||||
}
|
||||
tok, err := conf.TokenSource(ctx).Token()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("mint gcp token: %w", err)
|
||||
}
|
||||
return tok.AccessToken, nil
|
||||
}
|
||||
|
||||
// decorate turns raw vendor ids into the models the caller renders, attaching
|
||||
// the rates the request would actually be billed at.
|
||||
//
|
||||
// Rates come from the live default pricing table rather than the compiled-in
|
||||
// catalog, because that is the table the synthesiser ships to the proxy: an
|
||||
// operator running a defaults_llm_pricing.yaml would otherwise be shown one
|
||||
// price in the form and charged another. It is also the same lookup the catalog
|
||||
// endpoint prefills from, so a model reached by either route prices identically.
|
||||
func decorate(entry catalog.Provider, ids []listedModel) []Model {
|
||||
out := make([]Model, 0, len(ids))
|
||||
seen := make(map[string]struct{}, len(ids))
|
||||
for _, listed := range ids {
|
||||
if listed.id == "" {
|
||||
continue
|
||||
}
|
||||
if _, dup := seen[listed.id]; dup {
|
||||
continue
|
||||
}
|
||||
seen[listed.id] = struct{}{}
|
||||
|
||||
// The table keys pricing by the normalised id while the vendor issues
|
||||
// the wire form, so normalise before looking it up — otherwise every
|
||||
// Bedrock profile would report unpriced.
|
||||
model := Model{ID: listed.id, Label: listed.label}
|
||||
if rate, known := pricing.LookupDefault(entry.PricingSurfaces, normalizeForPricing(entry.ID, listed.id)); known {
|
||||
model.PricingKnown = true
|
||||
model.InputPer1k = rate.InputPer1k
|
||||
model.OutputPer1k = rate.OutputPer1k
|
||||
model.CachedInputPer1k = rate.CachedInputPer1k
|
||||
model.CacheReadPer1k = rate.CacheReadPer1k
|
||||
model.CacheCreationPer1k = rate.CacheCreationPer1k
|
||||
}
|
||||
out = append(out, model)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// refuseRedirect is the redirect policy every discovery request runs under. A
|
||||
// redirect is a way to move the request to a host checkPublicHost never saw,
|
||||
// so none are followed.
|
||||
func refuseRedirect(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
}
|
||||
|
||||
func (c *Client) httpClient() *http.Client {
|
||||
if c.HTTPClient != nil {
|
||||
if c.HTTPClient.CheckRedirect != nil {
|
||||
return c.HTTPClient
|
||||
}
|
||||
// An injected client that states no policy still gets ours: the
|
||||
// no-redirect guarantee should not depend on the caller remembering it.
|
||||
//
|
||||
// Copied rather than assigned into: one Client is shared by every
|
||||
// request for the process's lifetime, so writing to its fields here
|
||||
// would race across request goroutines. The copy shares the Transport,
|
||||
// which is safe for concurrent use by design.
|
||||
clone := *c.HTTPClient
|
||||
clone.CheckRedirect = refuseRedirect
|
||||
return &clone
|
||||
}
|
||||
transport := guardedTransport
|
||||
if c.AllowPrivateHosts {
|
||||
transport = http.DefaultTransport
|
||||
}
|
||||
return &http.Client{
|
||||
Timeout: fetchTimeout,
|
||||
Transport: transport,
|
||||
CheckRedirect: refuseRedirect,
|
||||
}
|
||||
}
|
||||
|
||||
// guardedTransport dials only addresses isPublic accepts.
|
||||
//
|
||||
// checkPublicHost resolves the host itself, and the transport then resolves it
|
||||
// again when it dials — two lookups of a name whose owner chooses the answers.
|
||||
// A record that returns a public address to the first and 127.0.0.1 to the
|
||||
// second passes the guard and reaches loopback anyway, which is the whole of
|
||||
// DNS rebinding. Re-checking at the socket closes that window: whatever the
|
||||
// second lookup returned is what Control is handed, and an address the guard
|
||||
// refuses never gets connected.
|
||||
//
|
||||
// Shared package-wide rather than built per Fetch so connections and their
|
||||
// pool survive between calls; the guard holds no state.
|
||||
var guardedTransport = newGuardedTransport()
|
||||
|
||||
func newGuardedTransport() http.RoundTripper {
|
||||
base, ok := http.DefaultTransport.(*http.Transport)
|
||||
if !ok {
|
||||
// Something replaced the default transport. Fall back to it rather
|
||||
// than dropping its behaviour, and rely on checkPublicHost alone.
|
||||
return http.DefaultTransport
|
||||
}
|
||||
// Cloned so proxy settings, TLS defaults and timeouts come from the
|
||||
// standard transport rather than being restated here.
|
||||
transport := base.Clone()
|
||||
dialer := &net.Dialer{
|
||||
Timeout: fetchTimeout,
|
||||
KeepAlive: 30 * time.Second,
|
||||
Control: func(_, address string, _ syscall.RawConn) error {
|
||||
return guardDialAddress(address)
|
||||
},
|
||||
}
|
||||
transport.DialContext = dialer.DialContext
|
||||
return transport
|
||||
}
|
||||
|
||||
// guardDialAddress refuses a resolved socket address the discovery client has
|
||||
// no business connecting to. Control hands it over post-resolution and
|
||||
// pre-connect, once per address the dialer tries, so a name with several A
|
||||
// records is checked at each one.
|
||||
func guardDialAddress(address string) error {
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return fmt.Errorf("discovery dial address %q is unreadable", address)
|
||||
}
|
||||
addr, err := netip.ParseAddr(host)
|
||||
if err != nil {
|
||||
// Control is documented to receive a resolved address; anything else
|
||||
// is a state we cannot vet, so it does not get dialled.
|
||||
return fmt.Errorf("discovery dial address %q is not an IP", host)
|
||||
}
|
||||
if !isPublic(addr) {
|
||||
return fmt.Errorf("discovery refused to dial non-public address %s", addr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,532 @@
|
||||
package modeldiscovery
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/pricing"
|
||||
)
|
||||
|
||||
// stubTransport answers every request with one canned response and records the
|
||||
// request it was given, so a test can assert on the URL and headers the client
|
||||
// built without a network round trip.
|
||||
type stubTransport struct {
|
||||
status int
|
||||
body string
|
||||
got *http.Request
|
||||
}
|
||||
|
||||
func (s *stubTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
s.got = req
|
||||
status := s.status
|
||||
if status == 0 {
|
||||
status = http.StatusOK
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: status,
|
||||
Body: io.NopCloser(strings.NewReader(s.body)),
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Request: req,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// newStubClient returns a client that never leaves the process. The host guard
|
||||
// is disabled because it would otherwise resolve the vendor's real name, which
|
||||
// would make these tests depend on DNS.
|
||||
func newStubClient(status int, body string) (*Client, *stubTransport) {
|
||||
tr := &stubTransport{status: status, body: body}
|
||||
return &Client{
|
||||
HTTPClient: &http.Client{Transport: tr},
|
||||
AllowPrivateHosts: true,
|
||||
}, tr
|
||||
}
|
||||
|
||||
// The payloads below are trimmed from what the vendors actually returned in
|
||||
// the discovery e2e, rather than invented, so a parser that only works against
|
||||
// an idealised shape fails here.
|
||||
|
||||
const openAIListing = `{"object":"list","data":[
|
||||
{"id":"gpt-4o-mini","object":"model","created":1721172741,"owned_by":"system"},
|
||||
{"id":"gpt-4o","object":"model","created":1715367049,"owned_by":"system"}
|
||||
]}`
|
||||
|
||||
const anthropicListing = `{"data":[
|
||||
{"type":"model","id":"claude-haiku-4-5-20251001","display_name":"Claude Haiku 4.5"},
|
||||
{"type":"model","id":"claude-sonnet-4-6","display_name":"Claude Sonnet 4.6"}
|
||||
],"has_more":false}`
|
||||
|
||||
const bedrockListing = `{"inferenceProfileSummaries":[
|
||||
{"inferenceProfileId":"eu.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
"inferenceProfileName":"EU Anthropic Claude Haiku 4.5","status":"ACTIVE","type":"SYSTEM_DEFINED"},
|
||||
{"inferenceProfileId":"global.cohere.embed-v4:0",
|
||||
"inferenceProfileName":"Global Cohere Embed v4","status":"ACTIVE","type":"SYSTEM_DEFINED"},
|
||||
{"inferenceProfileId":"eu.meta.llama3-2-1b-instruct-v1:0",
|
||||
"inferenceProfileName":"EU Meta Llama 3.2 1B","status":"INACTIVE","type":"SYSTEM_DEFINED"}
|
||||
]}`
|
||||
|
||||
const vertexListing = `{"publisherModels":[
|
||||
{"name":"publishers/anthropic/models/claude-3-opus","versionId":"20240229","launchStage":"GA"},
|
||||
{"name":"publishers/anthropic/models/claude-sonnet-4-5","versionId":"20250929","launchStage":"GA"}
|
||||
]}`
|
||||
|
||||
func TestFetchOpenAIListing(t *testing.T) {
|
||||
cl, tr := newStubClient(http.StatusOK, openAIListing)
|
||||
|
||||
models, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "openai_api",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
APIKey: "sk-test",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "https://api.openai.com/v1/models", tr.got.URL.String())
|
||||
assert.Equal(t, "Bearer sk-test", tr.got.Header.Get("Authorization"),
|
||||
"the credential must be injected through the catalog's auth template")
|
||||
assert.Equal(t, []string{"gpt-4o-mini", "gpt-4o"}, ids(models))
|
||||
for _, m := range models {
|
||||
assert.True(t, m.PricingKnown, "both models are in the shipped catalog: %s", m.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchAnthropicSendsTheVersionHeader(t *testing.T) {
|
||||
cl, tr := newStubClient(http.StatusOK, anthropicListing)
|
||||
|
||||
models, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "anthropic_api",
|
||||
UpstreamURL: "https://api.anthropic.com",
|
||||
APIKey: "sk-ant-test",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Anthropic rejects a request without the version header, so a listing
|
||||
// that reached us at all proves it was sent — but assert it, because the
|
||||
// failure mode otherwise only shows up against the live API.
|
||||
assert.Equal(t, "2023-06-01", tr.got.Header.Get("anthropic-version"))
|
||||
assert.Equal(t, "sk-ant-test", tr.got.Header.Get("x-api-key"),
|
||||
"Anthropic takes a bare key under its own header, not a Bearer token")
|
||||
assert.Equal(t, "limit=1000", tr.got.URL.RawQuery)
|
||||
|
||||
assert.Equal(t, []string{"claude-haiku-4-5-20251001", "claude-sonnet-4-6"}, ids(models))
|
||||
assert.Equal(t, "Claude Haiku 4.5", models[0].Label)
|
||||
}
|
||||
|
||||
func TestFetchBedrockUsesTheControlPlaneAndKeepsWireIDs(t *testing.T) {
|
||||
cl, tr := newStubClient(http.StatusOK, bedrockListing)
|
||||
|
||||
models, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "bedrock_api",
|
||||
// The record's upstream is the RUNTIME host, which does not serve
|
||||
// listings. The catalog's own discovery host must win over it.
|
||||
UpstreamURL: "https://bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
Region: "eu-central-1",
|
||||
APIKey: "aws-bearer",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, "https://bedrock.eu-central-1.amazonaws.com/inference-profiles",
|
||||
tr.got.URL.String(), "listings come from the control plane, not the runtime host")
|
||||
|
||||
// Region-prefixed ids verbatim: the prefix is what makes them invocable
|
||||
// and it cannot be reconstructed — global.* alongside eu.* is exactly the
|
||||
// case that defeats deriving it from the configured region.
|
||||
assert.Equal(t, []string{
|
||||
"eu.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
"global.cohere.embed-v4:0",
|
||||
}, ids(models), "an INACTIVE profile must not be offered")
|
||||
|
||||
assert.True(t, models[0].PricingKnown,
|
||||
"the catalog prices anthropic.claude-haiku-4-5, which this id normalises to")
|
||||
assert.False(t, models[1].PricingKnown,
|
||||
"cohere embed is not in the shipped Bedrock catalog, so the operator must price it")
|
||||
|
||||
// The rates travel with the model, so the form can prefill an editable row
|
||||
// rather than making the operator look every price up by hand.
|
||||
assert.Positive(t, models[0].InputPer1k, "a priced model must carry its input rate")
|
||||
assert.Positive(t, models[0].OutputPer1k, "a priced model must carry its output rate")
|
||||
// An unpriced model is offered at zero and flagged, not withheld: the
|
||||
// vendor says the credential can reach it.
|
||||
assert.Zero(t, models[1].InputPer1k)
|
||||
assert.Zero(t, models[1].OutputPer1k)
|
||||
}
|
||||
|
||||
// TestDiscoveredRatesMatchTheCatalogEndpoint pins the two prefill paths to one
|
||||
// table. The provider form fills a model row either from the catalog response
|
||||
// or from a discovery response, and an operator who switches between them must
|
||||
// not see the price change — both must equal what the proxy will bill.
|
||||
func TestDiscoveredRatesMatchTheCatalogEndpoint(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusOK, openAIListing)
|
||||
|
||||
models, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "openai_api",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
APIKey: "sk-test",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, models)
|
||||
|
||||
entry, ok := catalog.Lookup("openai_api")
|
||||
require.True(t, ok)
|
||||
|
||||
for _, m := range models {
|
||||
want, known := pricing.LookupDefault(entry.PricingSurfaces, m.ID)
|
||||
require.True(t, known, "%s should be priced by the default table", m.ID)
|
||||
assert.Equal(t, want.InputPer1k, m.InputPer1k, "input rate for %s", m.ID)
|
||||
assert.Equal(t, want.OutputPer1k, m.OutputPer1k, "output rate for %s", m.ID)
|
||||
assert.Equal(t, want.CachedInputPer1k, m.CachedInputPer1k, "cached-input rate for %s", m.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchVertexJoinsNameAndVersion(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusOK, vertexListing)
|
||||
|
||||
models, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "vertex_ai_api",
|
||||
UpstreamURL: "https://us-east5-aiplatform.googleapis.com",
|
||||
Region: "us-east5",
|
||||
APIKey: "ya29.test-token",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Vertex addresses a model as "<id>@<version>" on rawPredict, and splits
|
||||
// those across two fields in the listing.
|
||||
assert.Equal(t, []string{"claude-3-opus@20240229", "claude-sonnet-4-5@20250929"}, ids(models))
|
||||
assert.Equal(t, "claude-3-opus", models[0].Label)
|
||||
}
|
||||
|
||||
func TestFetchSurfacesTheVendorStatus(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusForbidden, `{"error":{"message":"no access"}}`)
|
||||
|
||||
_, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "openai_api",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
APIKey: "sk-test",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "403",
|
||||
"an operator whose key lacks access needs to see which status the vendor returned")
|
||||
}
|
||||
|
||||
func TestFetchRejectsAProviderWithoutDiscovery(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusOK, openAIListing)
|
||||
|
||||
_, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "litellm_proxy",
|
||||
UpstreamURL: "https://gateway.example.com",
|
||||
APIKey: "sk-test",
|
||||
})
|
||||
assert.ErrorIs(t, err, ErrNoDiscovery,
|
||||
"a gateway with no listing endpoint must be distinguishable from a failure, so the caller can fall back")
|
||||
}
|
||||
|
||||
func TestFetchRequiresACredential(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusOK, openAIListing)
|
||||
|
||||
_, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "openai_api",
|
||||
UpstreamURL: "https://api.openai.com",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "API key")
|
||||
}
|
||||
|
||||
func TestDiscoveryURLNeedsARegionWhenTheHostTemplatesOne(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusOK, bedrockListing)
|
||||
|
||||
// An upstream that matches no catalog template — a proxy in front of
|
||||
// Bedrock, say — leaves nothing to read the region from. Refusing beats
|
||||
// guessing: an unsubstituted placeholder would dial a host that does not
|
||||
// exist, and a guessed region would dial the wrong account's endpoint.
|
||||
_, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "bedrock_api",
|
||||
UpstreamURL: "https://bedrock.internal-proxy.example.com",
|
||||
APIKey: "aws-bearer",
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "region")
|
||||
}
|
||||
|
||||
// TestHostGuardRejectsNonPublicAddresses is the SSRF guard. Management holds a
|
||||
// credential for every provider, so an upstream pointed at an internal address
|
||||
// would turn discovery into a way to probe — and hand a token to — the
|
||||
// management server's own network.
|
||||
func TestHostGuardRejectsNonPublicAddresses(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
addr string
|
||||
want bool
|
||||
}{
|
||||
{"loopback v4", "127.0.0.1", false},
|
||||
{"loopback v6", "::1", false},
|
||||
{"private 10/8", "10.0.0.5", false},
|
||||
{"private 172.16/12", "172.16.4.1", false},
|
||||
{"private 192.168/16", "192.168.1.1", false},
|
||||
{"link-local", "169.254.169.254", false}, // cloud metadata
|
||||
{"unspecified", "0.0.0.0", false},
|
||||
{"multicast", "224.0.0.1", false},
|
||||
{"netbird overlay 100.64/10", "100.90.1.2", false},
|
||||
{"v4-mapped loopback", "::ffff:127.0.0.1", false},
|
||||
{"public v4", "1.1.1.1", true},
|
||||
{"public v6", "2606:4700:4700::1111", true},
|
||||
{"just outside CGNAT", "100.128.0.1", true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
addr, err := netip.ParseAddr(tc.addr)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tc.want, isPublic(addr))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestHostGuardResolvesAndRejectsLocalhost(t *testing.T) {
|
||||
cl := &Client{}
|
||||
err := cl.checkPublicHost("localhost")
|
||||
require.Error(t, err, "a name resolving to loopback must be refused, not just a literal address")
|
||||
assert.Contains(t, err.Error(), "non-public")
|
||||
}
|
||||
|
||||
// TestRedirectsAreNotFollowed covers a gap the other tests leave open: they all
|
||||
// inject an HTTPClient, which bypasses httpClient() and therefore the redirect
|
||||
// policy entirely. The policy is a security control — a 302 moves the request
|
||||
// to a host checkPublicHost never resolved — so it needs a test that goes
|
||||
// through the constructor the manager actually uses.
|
||||
func TestRedirectsAreNotFollowed(t *testing.T) {
|
||||
var hits int
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits++
|
||||
http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
for name, cl := range map[string]*Client{
|
||||
// The production shape: no injected client at all.
|
||||
"default client": {AllowPrivateHosts: true},
|
||||
// An injected client that states no policy must inherit ours rather
|
||||
// than silently chasing the redirect.
|
||||
"injected client with no policy": {
|
||||
AllowPrivateHosts: true,
|
||||
HTTPClient: &http.Client{},
|
||||
},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
hits = 0
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, err := cl.httpClient().Do(req)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = resp.Body.Close() })
|
||||
|
||||
assert.Equal(t, http.StatusFound, resp.StatusCode,
|
||||
"the redirect must be surfaced, not followed to an unchecked host")
|
||||
assert.Equal(t, 1, hits, "exactly one request must leave the client")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestInjectedClientKeepsItsOwnRedirectPolicy pins that the default above is a
|
||||
// default, not an override, and that supplying it does not mutate the caller's
|
||||
// client — one Client is shared across every request, so a write here would
|
||||
// race.
|
||||
func TestInjectedClientKeepsItsOwnRedirectPolicy(t *testing.T) {
|
||||
own := func(*http.Request, []*http.Request) error { return nil }
|
||||
injected := &http.Client{CheckRedirect: own}
|
||||
cl := &Client{HTTPClient: injected}
|
||||
|
||||
assert.Same(t, injected, cl.httpClient(),
|
||||
"a client that states a policy must be handed back untouched")
|
||||
|
||||
bare := &http.Client{}
|
||||
cl = &Client{HTTPClient: bare}
|
||||
require.NotSame(t, bare, cl.httpClient(), "the policy must be applied to a copy")
|
||||
assert.Nil(t, bare.CheckRedirect, "the caller's client must not be written to")
|
||||
}
|
||||
|
||||
// TestDialGuardRejectsRebindingToANonPublicAddress covers the window between
|
||||
// the two DNS lookups. checkPublicHost resolves the host, then the transport
|
||||
// resolves it again to dial; a name whose owner answers the first with a public
|
||||
// address and the second with 127.0.0.1 would otherwise pass the guard and
|
||||
// still reach loopback. The dial-time check sees whatever the second lookup
|
||||
// actually returned.
|
||||
func TestDialGuardRejectsRebindingToANonPublicAddress(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
address string
|
||||
wantErr string
|
||||
}{
|
||||
{"loopback", "127.0.0.1:443", "non-public"},
|
||||
{"cloud metadata", "169.254.169.254:80", "non-public"},
|
||||
{"rfc1918", "10.1.2.3:443", "non-public"},
|
||||
{"netbird overlay", "100.90.1.2:443", "non-public"},
|
||||
{"loopback v6", "[::1]:443", "non-public"},
|
||||
{"unresolved name", "evil.example.com:443", "not an IP"},
|
||||
{"no port", "1.1.1.1", "unreadable"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := guardDialAddress(tc.address)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tc.wantErr)
|
||||
})
|
||||
}
|
||||
|
||||
assert.NoError(t, guardDialAddress("1.1.1.1:443"), "a public address must still be dialled")
|
||||
assert.NoError(t, guardDialAddress("[2606:4700:4700::1111]:443"))
|
||||
}
|
||||
|
||||
// TestDialGuardIsInstalledOnTheDefaultClient pins the wiring rather than the
|
||||
// guard: a correct guard nothing calls protects nothing.
|
||||
func TestDialGuardIsInstalledOnTheDefaultClient(t *testing.T) {
|
||||
cl := &Client{}
|
||||
transport, ok := cl.httpClient().Transport.(*http.Transport)
|
||||
require.True(t, ok, "the default discovery client must carry the guarded transport")
|
||||
require.NotNil(t, transport.DialContext, "the guarded transport must dial through the guard")
|
||||
|
||||
_, err := transport.DialContext(context.Background(), "tcp", "127.0.0.1:9")
|
||||
require.Error(t, err, "the guard must refuse loopback even when the caller dials it directly")
|
||||
assert.Contains(t, err.Error(), "non-public")
|
||||
|
||||
// Tests point the client at a loopback server on purpose, so the opt-out
|
||||
// has to reach the dialer too.
|
||||
relaxed := &Client{AllowPrivateHosts: true}
|
||||
assert.Equal(t, http.DefaultTransport, relaxed.httpClient().Transport)
|
||||
}
|
||||
|
||||
// TestCallerInputFailuresAreMarkedInvalid keeps the handler's 400 mapping
|
||||
// honest: it branches on this sentinel, so an unmarked caller-input failure
|
||||
// silently becomes a 500.
|
||||
func TestCallerInputFailuresAreMarkedInvalid(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
req Request
|
||||
}{
|
||||
{"unknown provider", Request{CatalogID: "not_a_provider", APIKey: "k"}},
|
||||
{"unusable upstream", Request{CatalogID: "openai_api", UpstreamURL: "://", APIKey: "k"}},
|
||||
{"missing api key", Request{CatalogID: "openai_api", UpstreamURL: "https://api.openai.com"}},
|
||||
{"no region to read", Request{
|
||||
CatalogID: "bedrock_api",
|
||||
UpstreamURL: "https://bedrock-runtime.amazonaws.com",
|
||||
APIKey: "aws-bearer",
|
||||
}},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusOK, openAIListing)
|
||||
_, err := cl.Fetch(context.Background(), tc.req)
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, ErrInvalidRequest)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestEveryDiscoveryEntryHasAParser keeps the catalog and the parser table from
|
||||
// drifting: adding a Discovery block with a shape nothing parses would fail
|
||||
// only at runtime, in front of an operator.
|
||||
func TestEveryDiscoveryEntryHasAParser(t *testing.T) {
|
||||
for _, entry := range catalog.All() {
|
||||
if entry.Discovery == nil {
|
||||
continue
|
||||
}
|
||||
t.Run(entry.ID, func(t *testing.T) {
|
||||
assert.NotEmpty(t, entry.Discovery.Path, "a discovery entry needs a path")
|
||||
_, err := parseListing(entry.Discovery.Shape, []byte(`{}`))
|
||||
assert.NoError(t, err, "shape %q has no parser", entry.Discovery.Shape)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func ids(models []Model) []string {
|
||||
out := make([]string, 0, len(models))
|
||||
for _, m := range models {
|
||||
out = append(out, m.ID)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestRegionIsReadBackFromTheUpstream covers the reason the API takes no
|
||||
// region field: a provider record has none, and the operator already encoded
|
||||
// it in the upstream host when they configured inference.
|
||||
func TestRegionIsReadBackFromTheUpstream(t *testing.T) {
|
||||
cl, tr := newStubClient(http.StatusOK, bedrockListing)
|
||||
|
||||
_, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "bedrock_api",
|
||||
UpstreamURL: "https://bedrock-runtime.us-west-2.amazonaws.com",
|
||||
APIKey: "aws-bearer",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "bedrock.us-west-2.amazonaws.com", tr.got.URL.Host)
|
||||
}
|
||||
|
||||
func TestRegionFromUpstream(t *testing.T) {
|
||||
bedrock, ok := catalog.Lookup("bedrock_api")
|
||||
require.True(t, ok)
|
||||
vertex, ok := catalog.Lookup("vertex_ai_api")
|
||||
require.True(t, ok)
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
entry catalog.Provider
|
||||
upstream string
|
||||
want string
|
||||
}{
|
||||
{"bedrock runtime host", bedrock, "https://bedrock-runtime.eu-central-1.amazonaws.com", "eu-central-1"},
|
||||
{"bedrock without scheme", bedrock, "bedrock-runtime.ap-south-1.amazonaws.com", "ap-south-1"},
|
||||
{"vertex regional host", vertex, "https://us-east5-aiplatform.googleapis.com", "us-east5"},
|
||||
// A proxied or self-hosted upstream matches no template, and guessing
|
||||
// a region from it would build a URL pointing somewhere arbitrary.
|
||||
{"unrelated upstream", bedrock, "https://llm.internal.example.com", ""},
|
||||
{"vertex global host has no region segment", vertex, "https://aiplatform.googleapis.com", ""},
|
||||
// Bedrock's regionless endpoint carries both halves of the template at
|
||||
// once, with nothing between them. It has to read as "no region here"
|
||||
// rather than as an inverted slice range.
|
||||
{"bedrock regionless endpoint", bedrock, "https://bedrock-runtime.amazonaws.com", ""},
|
||||
{"bedrock regionless without scheme", bedrock, "bedrock-runtime.amazonaws.com", ""},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
assert.Equal(t, tc.want, RegionFromUpstream(tc.entry, tc.upstream))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// bedrockGeoListing carries profiles from geographies the original prefix list
|
||||
// did not name. Every one reduces to a catalog key, so every one must arrive
|
||||
// priced — an unstripped geography is what made a real account's listing come
|
||||
// back almost entirely at zero.
|
||||
const bedrockGeoListing = `{"inferenceProfileSummaries":[
|
||||
{"inferenceProfileId":"jp.anthropic.claude-sonnet-5-20260514-v1:0",
|
||||
"inferenceProfileName":"JP Anthropic Claude Sonnet 5","status":"ACTIVE","type":"SYSTEM_DEFINED"},
|
||||
{"inferenceProfileId":"au.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
"inferenceProfileName":"AU Anthropic Claude Haiku 4.5","status":"ACTIVE","type":"SYSTEM_DEFINED"},
|
||||
{"inferenceProfileId":"us-gov.anthropic.claude-sonnet-5-20260514-v1:0",
|
||||
"inferenceProfileName":"GovCloud Anthropic Claude Sonnet 5","status":"ACTIVE","type":"SYSTEM_DEFINED"}
|
||||
]}`
|
||||
|
||||
func TestBedrockProfilesFromAnyGeographyArrivePriced(t *testing.T) {
|
||||
cl, _ := newStubClient(http.StatusOK, bedrockGeoListing)
|
||||
|
||||
models, err := cl.Fetch(context.Background(), Request{
|
||||
CatalogID: "bedrock_api",
|
||||
UpstreamURL: "https://bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
APIKey: "aws-token",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, models, 3)
|
||||
|
||||
for _, m := range models {
|
||||
assert.True(t, m.PricingKnown, "%s must resolve to a catalog rate", m.ID)
|
||||
assert.Greater(t, m.InputPer1k, 0.0, "input rate for %s", m.ID)
|
||||
assert.Greater(t, m.OutputPer1k, 0.0, "output rate for %s", m.ID)
|
||||
assert.Greater(t, m.CacheReadPer1k, 0.0, "cache-read rate for %s", m.ID)
|
||||
}
|
||||
|
||||
// The wire id is preserved whatever the pricing key reduced to: it is the
|
||||
// only form that works at invoke time.
|
||||
assert.Equal(t, "jp.anthropic.claude-sonnet-5-20260514-v1:0", models[0].ID)
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package modeldiscovery
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
sharedllm "github.com/netbirdio/netbird/shared/llm"
|
||||
)
|
||||
|
||||
// listedModel is one entry lifted out of a vendor listing before the catalog
|
||||
// is consulted about it.
|
||||
type listedModel struct {
|
||||
id string
|
||||
label string
|
||||
}
|
||||
|
||||
// parseListing extracts model ids from a vendor listing. Each vendor invented
|
||||
// its own envelope, and the shape is declared by the catalog rather than
|
||||
// sniffed, so a vendor that changes shape fails loudly instead of silently
|
||||
// returning nothing.
|
||||
func parseListing(shape catalog.ListingShape, body []byte) ([]listedModel, error) {
|
||||
switch shape {
|
||||
case catalog.ShapeOpenAIData:
|
||||
return parseOpenAIData(body)
|
||||
case catalog.ShapeBedrockInferenceProfiles:
|
||||
return parseBedrockInferenceProfiles(body)
|
||||
case catalog.ShapeVertexPublisherModels:
|
||||
return parseVertexPublisherModels(body)
|
||||
default:
|
||||
return nil, fmt.Errorf("no parser for listing shape %q", shape)
|
||||
}
|
||||
}
|
||||
|
||||
// parseOpenAIData reads {"data":[{"id":…}]}, which OpenAI defined and
|
||||
// Anthropic adopted. Anthropic additionally supplies display_name.
|
||||
func parseOpenAIData(body []byte) ([]listedModel, error) {
|
||||
var doc struct {
|
||||
Data []struct {
|
||||
ID string `json:"id"`
|
||||
DisplayName string `json:"display_name"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &doc); err != nil {
|
||||
return nil, fmt.Errorf("decode model listing: %w", err)
|
||||
}
|
||||
out := make([]listedModel, 0, len(doc.Data))
|
||||
for _, entry := range doc.Data {
|
||||
out = append(out, listedModel{id: entry.ID, label: entry.DisplayName})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// parseBedrockInferenceProfiles reads
|
||||
// {"inferenceProfileSummaries":[{"inferenceProfileId":…}]}.
|
||||
//
|
||||
// The profile id is taken verbatim because its region prefix (eu., us.,
|
||||
// global.) is what makes it invocable, and it is not derivable from the
|
||||
// configured region — an account in one region legitimately holds global.*
|
||||
// profiles alongside its regional ones.
|
||||
//
|
||||
// Only ACTIVE profiles are offered: AWS reports others, and registering one
|
||||
// would produce a model that routes inside NetBird and fails at AWS.
|
||||
func parseBedrockInferenceProfiles(body []byte) ([]listedModel, error) {
|
||||
var doc struct {
|
||||
Summaries []struct {
|
||||
ID string `json:"inferenceProfileId"`
|
||||
Name string `json:"inferenceProfileName"`
|
||||
Status string `json:"status"`
|
||||
} `json:"inferenceProfileSummaries"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &doc); err != nil {
|
||||
return nil, fmt.Errorf("decode inference-profile listing: %w", err)
|
||||
}
|
||||
out := make([]listedModel, 0, len(doc.Summaries))
|
||||
for _, entry := range doc.Summaries {
|
||||
if entry.Status != "" && !strings.EqualFold(entry.Status, "ACTIVE") {
|
||||
continue
|
||||
}
|
||||
out = append(out, listedModel{id: entry.ID, label: entry.Name})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// parseVertexPublisherModels reads {"publisherModels":[{"name":…}]}, where
|
||||
// name is a resource path ("publishers/anthropic/models/claude-3-opus") and
|
||||
// the version lives in a separate field.
|
||||
//
|
||||
// Vertex addresses a model as "<id>@<version>" on the rawPredict path, so the
|
||||
// two are joined here: reporting the bare name would hand the operator an id
|
||||
// that looks usable and is not.
|
||||
func parseVertexPublisherModels(body []byte) ([]listedModel, error) {
|
||||
var doc struct {
|
||||
Models []struct {
|
||||
Name string `json:"name"`
|
||||
VersionID string `json:"versionId"`
|
||||
} `json:"publisherModels"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &doc); err != nil {
|
||||
return nil, fmt.Errorf("decode publisher-model listing: %w", err)
|
||||
}
|
||||
out := make([]listedModel, 0, len(doc.Models))
|
||||
for _, entry := range doc.Models {
|
||||
id := entry.Name
|
||||
if slash := strings.LastIndex(id, "/"); slash >= 0 {
|
||||
id = id[slash+1:]
|
||||
}
|
||||
if id == "" {
|
||||
continue
|
||||
}
|
||||
label := id
|
||||
if entry.VersionID != "" {
|
||||
id += "@" + entry.VersionID
|
||||
}
|
||||
out = append(out, listedModel{id: id, label: label})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// normalizeForPricing maps a vendor's wire id onto the key the catalog prices
|
||||
// it under. It mirrors the synthesiser's normalizePricingModelID: the two must
|
||||
// agree, or a model reported here as priced would meter at the default rate
|
||||
// instead of the operator's.
|
||||
func normalizeForPricing(catalogProviderID, modelID string) string {
|
||||
switch {
|
||||
case catalog.IsBedrockPathStyle(catalogProviderID):
|
||||
return sharedllm.NormalizeBedrockModel(modelID)
|
||||
case catalog.IsVertexPathStyle(catalogProviderID):
|
||||
return sharedllm.NormalizeVertexModel(modelID)
|
||||
default:
|
||||
return modelID
|
||||
}
|
||||
}
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -47,17 +47,11 @@ var supplementalDefaults = map[string]map[string]Entry{
|
||||
"gpt-5-nano": {InputPer1k: 0.00005, OutputPer1k: 0.0004, CachedInputPer1k: 0.000005},
|
||||
},
|
||||
"anthropic": {
|
||||
// claude-opus-5 is not yet in the catalog lineup but gateway /
|
||||
// grandfathered traffic uses it; priced so it isn't skipped.
|
||||
"claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
|
||||
// "kimi-k3[1m]" is the 1M-context alias some Claude Code guides
|
||||
// configure against Moonshot's Anthropic-compatible endpoint;
|
||||
// priced identically to kimi-k3 so those requests aren't skipped.
|
||||
"kimi-k3[1m]": {InputPer1k: 0.003, OutputPer1k: 0.015, CacheReadPer1k: 0.0003},
|
||||
},
|
||||
"bedrock": {
|
||||
"anthropic.claude-opus-5": {InputPer1k: 0.005, OutputPer1k: 0.025, CacheReadPer1k: 0.0005, CacheCreationPer1k: 0.00625},
|
||||
},
|
||||
}
|
||||
|
||||
var (
|
||||
|
||||
@@ -82,6 +82,11 @@ anthropic:
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
claude-sonnet-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
kimi-k3:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
@@ -145,6 +150,11 @@ bedrock:
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
anthropic.claude-sonnet-5:
|
||||
input_per_1k: 0.003
|
||||
output_per_1k: 0.015
|
||||
cache_read_per_1k: 0.0003
|
||||
cache_creation_per_1k: 0.00375
|
||||
meta.llama3-3-70b-instruct:
|
||||
input_per_1k: 0.00072
|
||||
output_per_1k: 0.00072
|
||||
|
||||
@@ -116,11 +116,13 @@ func TestDefaultTable_PinnedRates(t *testing.T) {
|
||||
assert.InDelta(t, 0.010, fable.InputPer1k, 1e-9, "claude-fable-5 input")
|
||||
assert.InDelta(t, 0.0125, fable.CacheCreationPer1k, 1e-9, "claude-fable-5 cache creation")
|
||||
|
||||
// Supplementals present on their surfaces.
|
||||
// Every id below must stay priced whichever source provides it: the
|
||||
// catalog lineup for the current Claude 5 family, supplementalDefaults
|
||||
// for the ids the dashboard deliberately doesn't offer.
|
||||
for surface, ids := range map[string][]string{
|
||||
"openai": {"gpt-5", "gpt-5-mini", "gpt-5-nano"},
|
||||
"anthropic": {"claude-opus-5", "kimi-k3[1m]", "kimi-k3"},
|
||||
"bedrock": {"anthropic.claude-opus-5"},
|
||||
"anthropic": {"claude-opus-5", "claude-sonnet-5", "kimi-k3[1m]", "kimi-k3"},
|
||||
"bedrock": {"anthropic.claude-opus-5", "anthropic.claude-sonnet-5"},
|
||||
} {
|
||||
for _, id := range ids {
|
||||
_, ok := table[surface][id]
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/modeldiscovery"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
|
||||
@@ -211,7 +212,19 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
|
||||
groupIndex := indexProviderGroups(enabledPolicies)
|
||||
|
||||
routerCfgJSON, err := buildRouterConfigJSON(enabledProviders, groupIndex)
|
||||
// The proxy guardrail is a per-provider fail-closed backstop; the
|
||||
// authoritative per-policy/group decision is management's
|
||||
// SelectPolicyForRequest. A provider lands in that map only when every
|
||||
// authorising policy restricts models.
|
||||
providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID)
|
||||
|
||||
// Discovery gets the finer view: per policy rather than flattened per
|
||||
// provider, so a listing can be bounded to what the calling groups may
|
||||
// actually use instead of the union across everyone who reaches the
|
||||
// provider.
|
||||
modelPolicies := buildModelPolicies(enabledPolicies, guardrailsByID)
|
||||
|
||||
routerCfgJSON, err := buildRouterConfigJSON(enabledProviders, groupIndex, modelPolicies)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -228,11 +241,6 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
|
||||
mergedGuardrails := mergeGuardrails(enabledPolicies, guardrailsByID)
|
||||
applyAccountCollectionControls(&mergedGuardrails, settings)
|
||||
// The proxy guardrail is a per-provider fail-closed backstop; the
|
||||
// authoritative per-policy/group decision is management's
|
||||
// SelectPolicyForRequest. A provider lands in this map only when every
|
||||
// authorising policy restricts models.
|
||||
providerAllowlists := buildProviderAllowlists(enabledPolicies, guardrailsByID)
|
||||
guardrailJSON, err := marshalGuardrailConfig(providerAllowlists, mergedGuardrails.PromptCapture)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -351,6 +359,11 @@ type routerProviderRoute struct {
|
||||
AuthHeaderName string `json:"auth_header_name"`
|
||||
AuthHeaderValue string `json:"auth_header_value"`
|
||||
AllowedGroupIDs []string `json:"allowed_group_ids,omitempty"`
|
||||
// ModelPolicies is one entry per enabled policy authorising this provider,
|
||||
// carrying that policy's source groups and the models it permits. The
|
||||
// router bounds a model listing with it, so a provider two groups reach
|
||||
// under different allowlists offers each only its own.
|
||||
ModelPolicies []routerModelPolicy `json:"model_policies,omitempty"`
|
||||
// Vertex marks a Google Vertex AI provider, whose requests carry the
|
||||
// model in the URL path. The router selects it by path, bypassing the
|
||||
// model/vendor table.
|
||||
@@ -368,6 +381,9 @@ type routerProviderRoute struct {
|
||||
// proxy dials this provider's upstream. For self-hosted / internal gateways
|
||||
// behind a private or self-signed certificate.
|
||||
SkipTLSVerify bool `json:"skip_tls_verify,omitempty"`
|
||||
// DiscoveryHost, when set, is the host serving this provider's model
|
||||
// listing, for a vendor that does not serve it from the inference host.
|
||||
DiscoveryHost string `json:"discovery_host,omitempty"`
|
||||
}
|
||||
|
||||
// indexProviderGroups walks the enabled policies and returns, per
|
||||
@@ -422,7 +438,7 @@ func indexProviderGroups(policies []*types.Policy) map[string][]string {
|
||||
// path-prefix tiebreak. Providers no enabled policy authorises
|
||||
// (orphans) are intentionally OMITTED so the router never observes a
|
||||
// route with an empty ACL.
|
||||
func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]string) ([]byte, error) {
|
||||
func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]string, modelPolicies map[string][]routerModelPolicy) ([]byte, error) {
|
||||
cfg := routerConfig{Providers: make([]routerProviderRoute, 0, len(providers))}
|
||||
for _, p := range providers {
|
||||
groups, hasPolicy := groupIndex[p.ID]
|
||||
@@ -435,6 +451,9 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("router config for provider %s: %w", p.ID, err)
|
||||
}
|
||||
// Lookup rather than assume: an unknown provider id yields the zero
|
||||
// entry, which declares no discovery and so contributes nothing.
|
||||
catalogEntry, _ := catalog.Lookup(p.ProviderID)
|
||||
headerName, headerValue, gcpSAKeyB64, err := providerAuthHeader(p)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -449,10 +468,12 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]
|
||||
AuthHeaderName: headerName,
|
||||
AuthHeaderValue: headerValue,
|
||||
AllowedGroupIDs: groups,
|
||||
ModelPolicies: modelPolicies[p.ID],
|
||||
Vertex: catalog.IsVertexPathStyle(p.ProviderID),
|
||||
Bedrock: catalog.IsBedrockPathStyle(p.ProviderID),
|
||||
GCPServiceAccountKeyB64: gcpSAKeyB64,
|
||||
SkipTLSVerify: p.SkipTLSVerification,
|
||||
DiscoveryHost: discoveryHost(catalogEntry, p.UpstreamURL),
|
||||
})
|
||||
}
|
||||
out, err := json.Marshal(cfg)
|
||||
@@ -462,6 +483,33 @@ func buildRouterConfigJSON(providers []*types.Provider, groupIndex map[string][]
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// discoveryHost returns the host serving this provider's model listing when it
|
||||
// differs from the inference host, and empty when the two are the same — which
|
||||
// is true of every vendor but Bedrock, whose ListInferenceProfiles is a control
|
||||
// plane operation on bedrock.<region> while InvokeModel must go to
|
||||
// bedrock-runtime.<region>. One provider record therefore needs two hosts.
|
||||
//
|
||||
// The catalog declares the listing host; the region is recovered from the
|
||||
// upstream the operator configured, since a provider record carries no region
|
||||
// field. An upstream matching no catalog template yields empty rather than a
|
||||
// guess: a proxied or self-hosted Bedrock endpoint may serve both from one
|
||||
// place, and inventing a host would send the credential somewhere the operator
|
||||
// never configured.
|
||||
func discoveryHost(entry catalog.Provider, upstreamURL string) string {
|
||||
if entry.Discovery == nil || entry.Discovery.Host == "" {
|
||||
return ""
|
||||
}
|
||||
host := entry.Discovery.Host
|
||||
if !strings.Contains(host, catalog.RegionPlaceholder) {
|
||||
return host
|
||||
}
|
||||
region := modeldiscovery.RegionFromUpstream(entry, upstreamURL)
|
||||
if region == "" {
|
||||
return ""
|
||||
}
|
||||
return strings.ReplaceAll(host, catalog.RegionPlaceholder, region)
|
||||
}
|
||||
|
||||
// providerVendor returns the parser surface ("openai", "anthropic", …)
|
||||
// the provider speaks, sourced from its catalog entry's ParserID. The
|
||||
// router uses it to keep a request the parser tagged with a vendor on a
|
||||
@@ -1098,3 +1146,46 @@ func mergeGuardrail(g *types.Guardrail, merged *MergedGuardrails) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// routerModelPolicy mirrors the router's ModelPolicyRule: one authorising
|
||||
// policy's source groups plus the models it permits. Models is nil for a
|
||||
// policy that sets no model allowlist, which lifts the restriction for the
|
||||
// groups it binds — so nil and empty must survive the round trip distinctly.
|
||||
type routerModelPolicy struct {
|
||||
GroupIDs []string `json:"group_ids"`
|
||||
Models []string `json:"models"`
|
||||
}
|
||||
|
||||
// buildModelPolicies indexes, per provider, one rule for each enabled policy
|
||||
// authorising it: the policy's source groups and the models its guardrail
|
||||
// permits.
|
||||
//
|
||||
// This is deliberately finer than buildProviderAllowlists, which flattens the
|
||||
// same inputs into one list per provider for the proxy's fail-closed guardrail.
|
||||
// A flattened list cannot answer "what may THIS caller see", so a provider two
|
||||
// teams reach under different allowlists would offer each team the other's
|
||||
// models — a picker full of entries the next request refuses. Keeping the
|
||||
// source groups alongside the models lets the router answer it at request time,
|
||||
// where it knows the caller's groups.
|
||||
func buildModelPolicies(policies []*types.Policy, byID map[string]*types.Guardrail) map[string][]routerModelPolicy {
|
||||
out := make(map[string][]routerModelPolicy)
|
||||
for _, p := range policies {
|
||||
if p == nil || len(p.SourceGroups) == 0 {
|
||||
continue
|
||||
}
|
||||
restricted, models := policyModelAllowlist(p, byID)
|
||||
rule := routerModelPolicy{GroupIDs: append([]string(nil), p.SourceGroups...)}
|
||||
if restricted {
|
||||
// Never nil when restricted: an allowlist permitting nothing must
|
||||
// stay distinguishable from no allowlist at all.
|
||||
rule.Models = append([]string{}, models...)
|
||||
}
|
||||
for _, providerID := range p.DestinationProviderIDs {
|
||||
if providerID == "" {
|
||||
continue
|
||||
}
|
||||
out[providerID] = append(out[providerID], rule)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
@@ -103,3 +103,37 @@ func TestBuildCostMeterConfig_OrphanAndGatewayProviders(t *testing.T) {
|
||||
assert.NotContains(t, cfg.Pricing.Providers, "prov-litellm", "empty-models gateway needs no per-record entry")
|
||||
assert.NotEmpty(t, cfg.Pricing.Defaults["openai"], "defaults still ship so the gateway's catalog-model traffic is priced")
|
||||
}
|
||||
|
||||
// TestBuildCostMeterConfig_BedrockGeographyOutsideTheOriginalFour is the
|
||||
// accounting half of the geography bug. The docs tell operators to register a
|
||||
// Bedrock id exactly as AWS issues it, region prefix included, and the cost
|
||||
// meter keys its table by the normalized form. While the geography was matched
|
||||
// against a list of four, a profile issued anywhere else kept its prefix,
|
||||
// missed the catalog entry it was meant to inherit from, and billed with a
|
||||
// zero entry underneath the operator's own rates — so every cache bucket
|
||||
// metered free and a model priced only by catalog defaults metered at nothing
|
||||
// at all.
|
||||
func TestBuildCostMeterConfig_BedrockGeographyOutsideTheOriginalFour(t *testing.T) {
|
||||
for _, geo := range []string{"jp", "au", "ca", "sa", "us-gov"} {
|
||||
t.Run(geo, func(t *testing.T) {
|
||||
bedrock := &types.Provider{
|
||||
ID: "prov-bedrock",
|
||||
ProviderID: "bedrock_api",
|
||||
Enabled: true,
|
||||
Models: []types.ProviderModel{
|
||||
{ID: geo + ".anthropic.claude-sonnet-5-20260514-v1:0", InputPer1k: 0.003, OutputPer1k: 0.015},
|
||||
},
|
||||
}
|
||||
raw, err := buildCostMeterConfigJSON([]*types.Provider{bedrock}, map[string][]string{"prov-bedrock": {"grp"}})
|
||||
require.NoError(t, err)
|
||||
cfg := decodeCostMeterConfig(t, raw)
|
||||
|
||||
e, ok := cfg.Pricing.Providers["prov-bedrock"]["anthropic.claude-sonnet-5"]
|
||||
require.True(t, ok, "a %s profile must key by the same normalized id the parser emits", geo)
|
||||
assert.InDelta(t, 0.0003, e.CacheReadPer1k, 1e-9,
|
||||
"cache read must be inherited from the bedrock default entry, not left at zero")
|
||||
assert.InDelta(t, 0.00375, e.CacheCreationPer1k, 1e-9,
|
||||
"cache creation must be inherited from the bedrock default entry, not left at zero")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
)
|
||||
@@ -93,3 +94,75 @@ func TestBuildProviderAllowlists(t *testing.T) {
|
||||
"an enabled-but-empty allowlist is restricted with an empty set, not unrestricted")
|
||||
})
|
||||
}
|
||||
|
||||
// policyForGroups builds an enabled policy binding the given source groups to
|
||||
// the given providers under an optional guardrail.
|
||||
func policyForGroups(id string, groups []string, guardrailIDs []string, providerIDs ...string) *types.Policy {
|
||||
return &types.Policy{
|
||||
ID: id,
|
||||
Enabled: true,
|
||||
SourceGroups: groups,
|
||||
DestinationProviderIDs: providerIDs,
|
||||
GuardrailIDs: guardrailIDs,
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildModelPolicies covers the finer index discovery needs. Where
|
||||
// buildProviderAllowlists flattens every authorising policy into one list per
|
||||
// provider — enough for a fail-closed backstop, but blind to who is asking —
|
||||
// this keeps each policy's source groups beside its models so the router can
|
||||
// bound a listing to the calling groups.
|
||||
func TestBuildModelPolicies(t *testing.T) {
|
||||
byID := map[string]*types.Guardrail{
|
||||
"g-4o": allowlistGuardrail("g-4o", "acc-1", "gpt-4o"),
|
||||
"g-opus": allowlistGuardrail("g-opus", "acc-1", "claude-opus-4"),
|
||||
"g-disabled": {ID: "g-disabled", Checks: types.GuardrailChecks{ModelAllowlist: types.GuardrailModelAllowlist{Enabled: false, Models: []string{"gpt-4o"}}}},
|
||||
}
|
||||
|
||||
t.Run("each policy keeps its own groups and models", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForGroups("p1", []string{"grp-eng"}, []string{"g-4o"}, "prov-x"),
|
||||
policyForGroups("p2", []string{"grp-sales"}, []string{"g-opus"}, "prov-x"),
|
||||
}
|
||||
got := buildModelPolicies(policies, byID)
|
||||
assert.Equal(t, []routerModelPolicy{
|
||||
{GroupIDs: []string{"grp-eng"}, Models: []string{"gpt-4o"}},
|
||||
{GroupIDs: []string{"grp-sales"}, Models: []string{"claude-opus-4"}},
|
||||
}, got["prov-x"],
|
||||
"the two policies must stay separable so neither group is offered the other's models")
|
||||
})
|
||||
|
||||
t.Run("an unrestricted policy carries nil models", func(t *testing.T) {
|
||||
policies := []*types.Policy{
|
||||
policyForGroups("p1", []string{"grp-eng"}, []string{"g-4o"}, "prov-x"),
|
||||
policyForGroups("p2", []string{"grp-admin"}, nil, "prov-x"),
|
||||
}
|
||||
got := buildModelPolicies(policies, byID)
|
||||
assert.Nil(t, got["prov-x"][1].Models,
|
||||
"no allowlist must reach the router as nil, which lifts the restriction for its groups")
|
||||
})
|
||||
|
||||
t.Run("a disabled allowlist is not a restriction", func(t *testing.T) {
|
||||
policies := []*types.Policy{policyForGroups("p1", []string{"grp-eng"}, []string{"g-disabled"}, "prov-x")}
|
||||
got := buildModelPolicies(policies, byID)
|
||||
assert.Nil(t, got["prov-x"][0].Models,
|
||||
"a guardrail with the allowlist check off restricts nothing")
|
||||
})
|
||||
|
||||
t.Run("an enabled allowlist with no models permits nothing", func(t *testing.T) {
|
||||
byIDEmpty := map[string]*types.Guardrail{
|
||||
"g-empty": {ID: "g-empty", Checks: types.GuardrailChecks{ModelAllowlist: types.GuardrailModelAllowlist{Enabled: true}}},
|
||||
}
|
||||
policies := []*types.Policy{policyForGroups("p1", []string{"grp-eng"}, []string{"g-empty"}, "prov-x")}
|
||||
got := buildModelPolicies(policies, byIDEmpty)
|
||||
require.NotNil(t, got["prov-x"][0].Models,
|
||||
"an empty allowlist must not arrive as nil — that would read as unrestricted")
|
||||
assert.Empty(t, got["prov-x"][0].Models)
|
||||
})
|
||||
|
||||
t.Run("a policy binding no groups is skipped", func(t *testing.T) {
|
||||
policies := []*types.Policy{policyForGroups("p1", nil, []string{"g-4o"}, "prov-x")}
|
||||
assert.Empty(t, buildModelPolicies(policies, byID),
|
||||
"a policy with no source groups authorises nobody, so it bounds nobody's listing")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -6,10 +6,11 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
@@ -1245,3 +1246,57 @@ func TestSynthesizeServices_EmptyAPIKey_FailsClosed(t *testing.T) {
|
||||
require.Error(t, err, "synthesis must refuse a provider with no api key")
|
||||
assert.Contains(t, err.Error(), "no api key", "error must surface the missing credential")
|
||||
}
|
||||
|
||||
// TestDiscoveryHost pins which providers get a separate listing host. Getting
|
||||
// this wrong in either direction is costly: a missing host leaves Bedrock
|
||||
// discovery 404ing at AWS, and a host on the wrong provider would send that
|
||||
// provider's listing — and its credential — somewhere the operator never
|
||||
// configured.
|
||||
func TestDiscoveryHost(t *testing.T) {
|
||||
entry := func(id string) catalog.Provider {
|
||||
p, ok := catalog.Lookup(id)
|
||||
require.True(t, ok, "catalog entry %s must exist", id)
|
||||
return p
|
||||
}
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
entry catalog.Provider
|
||||
upstream string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
// ListInferenceProfiles is a control-plane operation; the runtime
|
||||
// host answers <UnknownOperationException/> for it.
|
||||
name: "bedrock splits the listing off the runtime host",
|
||||
entry: entry("bedrock_api"), upstream: "https://bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
want: "bedrock.eu-central-1.amazonaws.com",
|
||||
},
|
||||
{
|
||||
name: "bedrock in another region",
|
||||
entry: entry("bedrock_api"), upstream: "https://bedrock-runtime.us-west-2.amazonaws.com",
|
||||
want: "bedrock.us-west-2.amazonaws.com",
|
||||
},
|
||||
{
|
||||
// A proxied Bedrock endpoint may well serve both from one place,
|
||||
// and there is no region to read back out of it.
|
||||
name: "proxied bedrock upstream yields no discovery host",
|
||||
entry: entry("bedrock_api"), upstream: "https://bedrock.internal.example.com",
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "openai serves its listing from the same host",
|
||||
entry: entry("openai_api"), upstream: "https://api.openai.com",
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "vertex serves its listing from the same host",
|
||||
entry: entry("vertex_ai_api"), upstream: "https://us-east5-aiplatform.googleapis.com",
|
||||
want: "",
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
assert.Equal(t, tc.want, discoveryHost(tc.entry, tc.upstream))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package peers
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./manager.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package peers is a generated GoMock package.
|
||||
package peers
|
||||
@@ -9,18 +14,19 @@ import (
|
||||
net "net"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
network_map "github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
account "github.com/netbirdio/netbird/management/server/account"
|
||||
integrated_validator "github.com/netbirdio/netbird/management/server/integrations/integrated_validator"
|
||||
peer "github.com/netbirdio/netbird/management/server/peer"
|
||||
types "github.com/netbirdio/netbird/management/server/types"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -49,7 +55,7 @@ func (m *MockManager) CreateProxyPeer(ctx context.Context, accountID, peerKey, c
|
||||
}
|
||||
|
||||
// CreateProxyPeer indicates an expected call of CreateProxyPeer.
|
||||
func (mr *MockManagerMockRecorder) CreateProxyPeer(ctx, accountID, peerKey, cluster interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CreateProxyPeer(ctx, accountID, peerKey, cluster any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateProxyPeer", reflect.TypeOf((*MockManager)(nil).CreateProxyPeer), ctx, accountID, peerKey, cluster)
|
||||
}
|
||||
@@ -63,7 +69,7 @@ func (m *MockManager) DeletePeers(ctx context.Context, accountID string, peerIDs
|
||||
}
|
||||
|
||||
// DeletePeers indicates an expected call of DeletePeers.
|
||||
func (mr *MockManagerMockRecorder) DeletePeers(ctx, accountID, peerIDs, userID, checkConnected interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeletePeers(ctx, accountID, peerIDs, userID, checkConnected any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeletePeers", reflect.TypeOf((*MockManager)(nil).DeletePeers), ctx, accountID, peerIDs, userID, checkConnected)
|
||||
}
|
||||
@@ -78,7 +84,7 @@ func (m *MockManager) GetAllPeers(ctx context.Context, accountID, userID string)
|
||||
}
|
||||
|
||||
// GetAllPeers indicates an expected call of GetAllPeers.
|
||||
func (mr *MockManagerMockRecorder) GetAllPeers(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAllPeers(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllPeers", reflect.TypeOf((*MockManager)(nil).GetAllPeers), ctx, accountID, userID)
|
||||
}
|
||||
@@ -93,7 +99,7 @@ func (m *MockManager) GetPeer(ctx context.Context, accountID, userID, peerID str
|
||||
}
|
||||
|
||||
// GetPeer indicates an expected call of GetPeer.
|
||||
func (mr *MockManagerMockRecorder) GetPeer(ctx, accountID, userID, peerID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeer(ctx, accountID, userID, peerID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeer", reflect.TypeOf((*MockManager)(nil).GetPeer), ctx, accountID, userID, peerID)
|
||||
}
|
||||
@@ -108,7 +114,7 @@ func (m *MockManager) GetPeerAccountID(ctx context.Context, peerID string) (stri
|
||||
}
|
||||
|
||||
// GetPeerAccountID indicates an expected call of GetPeerAccountID.
|
||||
func (mr *MockManagerMockRecorder) GetPeerAccountID(ctx, peerID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerAccountID(ctx, peerID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerAccountID", reflect.TypeOf((*MockManager)(nil).GetPeerAccountID), ctx, peerID)
|
||||
}
|
||||
@@ -123,7 +129,7 @@ func (m *MockManager) GetPeerByTunnelIP(ctx context.Context, accountID string, i
|
||||
}
|
||||
|
||||
// GetPeerByTunnelIP indicates an expected call of GetPeerByTunnelIP.
|
||||
func (mr *MockManagerMockRecorder) GetPeerByTunnelIP(ctx, accountID, ip interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerByTunnelIP(ctx, accountID, ip any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerByTunnelIP", reflect.TypeOf((*MockManager)(nil).GetPeerByTunnelIP), ctx, accountID, ip)
|
||||
}
|
||||
@@ -138,7 +144,7 @@ func (m *MockManager) GetPeerID(ctx context.Context, peerKey string) (string, er
|
||||
}
|
||||
|
||||
// GetPeerID indicates an expected call of GetPeerID.
|
||||
func (mr *MockManagerMockRecorder) GetPeerID(ctx, peerKey interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerID(ctx, peerKey any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerID", reflect.TypeOf((*MockManager)(nil).GetPeerID), ctx, peerKey)
|
||||
}
|
||||
@@ -154,7 +160,7 @@ func (m *MockManager) GetPeerWithGroups(ctx context.Context, accountID, peerID s
|
||||
}
|
||||
|
||||
// GetPeerWithGroups indicates an expected call of GetPeerWithGroups.
|
||||
func (mr *MockManagerMockRecorder) GetPeerWithGroups(ctx, accountID, peerID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerWithGroups(ctx, accountID, peerID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerWithGroups", reflect.TypeOf((*MockManager)(nil).GetPeerWithGroups), ctx, accountID, peerID)
|
||||
}
|
||||
@@ -169,7 +175,7 @@ func (m *MockManager) GetPeersByGroupIDs(ctx context.Context, accountID string,
|
||||
}
|
||||
|
||||
// GetPeersByGroupIDs indicates an expected call of GetPeersByGroupIDs.
|
||||
func (mr *MockManagerMockRecorder) GetPeersByGroupIDs(ctx, accountID, groupsIDs interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeersByGroupIDs(ctx, accountID, groupsIDs any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeersByGroupIDs", reflect.TypeOf((*MockManager)(nil).GetPeersByGroupIDs), ctx, accountID, groupsIDs)
|
||||
}
|
||||
@@ -181,7 +187,7 @@ func (m *MockManager) SetAccountManager(accountManager account.Manager) {
|
||||
}
|
||||
|
||||
// SetAccountManager indicates an expected call of SetAccountManager.
|
||||
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetAccountManager", reflect.TypeOf((*MockManager)(nil).SetAccountManager), accountManager)
|
||||
}
|
||||
@@ -193,7 +199,7 @@ func (m *MockManager) SetIntegratedPeerValidator(integratedPeerValidator integra
|
||||
}
|
||||
|
||||
// SetIntegratedPeerValidator indicates an expected call of SetIntegratedPeerValidator.
|
||||
func (mr *MockManagerMockRecorder) SetIntegratedPeerValidator(integratedPeerValidator interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetIntegratedPeerValidator(integratedPeerValidator any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetIntegratedPeerValidator", reflect.TypeOf((*MockManager)(nil).SetIntegratedPeerValidator), integratedPeerValidator)
|
||||
}
|
||||
@@ -205,7 +211,7 @@ func (m *MockManager) SetNetworkMapController(networkMapController network_map.C
|
||||
}
|
||||
|
||||
// SetNetworkMapController indicates an expected call of SetNetworkMapController.
|
||||
func (mr *MockManagerMockRecorder) SetNetworkMapController(networkMapController interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetNetworkMapController(networkMapController any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetNetworkMapController", reflect.TypeOf((*MockManager)(nil).SetNetworkMapController), networkMapController)
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package proxy
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./manager.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package proxy is a generated GoMock package.
|
||||
package proxy
|
||||
@@ -9,14 +14,15 @@ import (
|
||||
reflect "reflect"
|
||||
time "time"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
proto "github.com/netbirdio/netbird/shared/management/proto"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -45,7 +51,7 @@ func (m *MockManager) CleanupStale(ctx context.Context, inactivityDuration time.
|
||||
}
|
||||
|
||||
// CleanupStale indicates an expected call of CleanupStale.
|
||||
func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration)
|
||||
}
|
||||
@@ -59,7 +65,7 @@ func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr s
|
||||
}
|
||||
|
||||
// ClusterRequireSubdomain indicates an expected call of ClusterRequireSubdomain.
|
||||
func (mr *MockManagerMockRecorder) ClusterRequireSubdomain(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterRequireSubdomain(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterRequireSubdomain", reflect.TypeOf((*MockManager)(nil).ClusterRequireSubdomain), ctx, clusterAddr)
|
||||
}
|
||||
@@ -73,7 +79,7 @@ func (m *MockManager) ClusterSupportsAppSec(ctx context.Context, clusterAddr str
|
||||
}
|
||||
|
||||
// ClusterSupportsAppSec indicates an expected call of ClusterSupportsAppSec.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsAppSec(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsAppSec(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsAppSec", reflect.TypeOf((*MockManager)(nil).ClusterSupportsAppSec), ctx, clusterAddr)
|
||||
}
|
||||
@@ -87,7 +93,7 @@ func (m *MockManager) ClusterSupportsCrowdSec(ctx context.Context, clusterAddr s
|
||||
}
|
||||
|
||||
// ClusterSupportsCrowdSec indicates an expected call of ClusterSupportsCrowdSec.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCrowdSec(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCrowdSec(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsCrowdSec", reflect.TypeOf((*MockManager)(nil).ClusterSupportsCrowdSec), ctx, clusterAddr)
|
||||
}
|
||||
@@ -101,7 +107,7 @@ func (m *MockManager) ClusterSupportsCustomPorts(ctx context.Context, clusterAdd
|
||||
}
|
||||
|
||||
// ClusterSupportsCustomPorts indicates an expected call of ClusterSupportsCustomPorts.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCustomPorts(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCustomPorts(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsCustomPorts", reflect.TypeOf((*MockManager)(nil).ClusterSupportsCustomPorts), ctx, clusterAddr)
|
||||
}
|
||||
@@ -115,7 +121,7 @@ func (m *MockManager) ClusterSupportsPrivate(ctx context.Context, clusterAddr st
|
||||
}
|
||||
|
||||
// ClusterSupportsPrivate indicates an expected call of ClusterSupportsPrivate.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
|
||||
}
|
||||
@@ -130,7 +136,7 @@ func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAd
|
||||
}
|
||||
|
||||
// Connect indicates an expected call of Connect.
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
|
||||
}
|
||||
@@ -145,7 +151,7 @@ func (m *MockManager) CountAccountProxies(ctx context.Context, accountID string)
|
||||
}
|
||||
|
||||
// CountAccountProxies indicates an expected call of CountAccountProxies.
|
||||
func (mr *MockManagerMockRecorder) CountAccountProxies(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CountAccountProxies(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountAccountProxies", reflect.TypeOf((*MockManager)(nil).CountAccountProxies), ctx, accountID)
|
||||
}
|
||||
@@ -159,7 +165,7 @@ func (m *MockManager) DeleteAccountCluster(ctx context.Context, clusterAddress,
|
||||
}
|
||||
|
||||
// DeleteAccountCluster indicates an expected call of DeleteAccountCluster.
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, clusterAddress, accountID)
|
||||
}
|
||||
@@ -173,7 +179,7 @@ func (m *MockManager) Disconnect(ctx context.Context, proxyID, sessionID string)
|
||||
}
|
||||
|
||||
// Disconnect indicates an expected call of Disconnect.
|
||||
func (mr *MockManagerMockRecorder) Disconnect(ctx, proxyID, sessionID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Disconnect(ctx, proxyID, sessionID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Disconnect", reflect.TypeOf((*MockManager)(nil).Disconnect), ctx, proxyID, sessionID)
|
||||
}
|
||||
@@ -188,7 +194,7 @@ func (m *MockManager) GetAccountProxy(ctx context.Context, accountID string) (*P
|
||||
}
|
||||
|
||||
// GetAccountProxy indicates an expected call of GetAccountProxy.
|
||||
func (mr *MockManagerMockRecorder) GetAccountProxy(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAccountProxy(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountProxy", reflect.TypeOf((*MockManager)(nil).GetAccountProxy), ctx, accountID)
|
||||
}
|
||||
@@ -203,7 +209,7 @@ func (m *MockManager) GetActiveClusterAddresses(ctx context.Context) ([]string,
|
||||
}
|
||||
|
||||
// GetActiveClusterAddresses indicates an expected call of GetActiveClusterAddresses.
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddresses(ctx interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddresses(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveClusterAddresses", reflect.TypeOf((*MockManager)(nil).GetActiveClusterAddresses), ctx)
|
||||
}
|
||||
@@ -218,7 +224,7 @@ func (m *MockManager) GetActiveClusterAddressesForAccount(ctx context.Context, a
|
||||
}
|
||||
|
||||
// GetActiveClusterAddressesForAccount indicates an expected call of GetActiveClusterAddressesForAccount.
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddressesForAccount(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddressesForAccount(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveClusterAddressesForAccount", reflect.TypeOf((*MockManager)(nil).GetActiveClusterAddressesForAccount), ctx, accountID)
|
||||
}
|
||||
@@ -232,7 +238,7 @@ func (m *MockManager) Heartbeat(ctx context.Context, p *Proxy) error {
|
||||
}
|
||||
|
||||
// Heartbeat indicates an expected call of Heartbeat.
|
||||
func (mr *MockManagerMockRecorder) Heartbeat(ctx, p interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Heartbeat(ctx, p any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Heartbeat", reflect.TypeOf((*MockManager)(nil).Heartbeat), ctx, p)
|
||||
}
|
||||
@@ -247,7 +253,7 @@ func (m *MockManager) IsClusterAddressAvailable(ctx context.Context, clusterAddr
|
||||
}
|
||||
|
||||
// IsClusterAddressAvailable indicates an expected call of IsClusterAddressAvailable.
|
||||
func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "IsClusterAddressAvailable", reflect.TypeOf((*MockManager)(nil).IsClusterAddressAvailable), ctx, clusterAddress, accountID)
|
||||
}
|
||||
@@ -256,6 +262,7 @@ func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress
|
||||
type MockController struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockControllerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockControllerMockRecorder is the mock recorder for MockController.
|
||||
@@ -298,7 +305,7 @@ func (m *MockController) GetProxiesForCluster(clusterAddr string) []string {
|
||||
}
|
||||
|
||||
// GetProxiesForCluster indicates an expected call of GetProxiesForCluster.
|
||||
func (mr *MockControllerMockRecorder) GetProxiesForCluster(clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) GetProxiesForCluster(clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetProxiesForCluster", reflect.TypeOf((*MockController)(nil).GetProxiesForCluster), clusterAddr)
|
||||
}
|
||||
@@ -312,7 +319,7 @@ func (m *MockController) RegisterProxyToCluster(ctx context.Context, clusterAddr
|
||||
}
|
||||
|
||||
// RegisterProxyToCluster indicates an expected call of RegisterProxyToCluster.
|
||||
func (mr *MockControllerMockRecorder) RegisterProxyToCluster(ctx, clusterAddr, proxyID interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) RegisterProxyToCluster(ctx, clusterAddr, proxyID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterProxyToCluster", reflect.TypeOf((*MockController)(nil).RegisterProxyToCluster), ctx, clusterAddr, proxyID)
|
||||
}
|
||||
@@ -324,7 +331,7 @@ func (m *MockController) SendServiceUpdateToCluster(ctx context.Context, account
|
||||
}
|
||||
|
||||
// SendServiceUpdateToCluster indicates an expected call of SendServiceUpdateToCluster.
|
||||
func (mr *MockControllerMockRecorder) SendServiceUpdateToCluster(ctx, accountID, update, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) SendServiceUpdateToCluster(ctx, accountID, update, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SendServiceUpdateToCluster", reflect.TypeOf((*MockController)(nil).SendServiceUpdateToCluster), ctx, accountID, update, clusterAddr)
|
||||
}
|
||||
@@ -338,7 +345,7 @@ func (m *MockController) UnregisterProxyFromCluster(ctx context.Context, cluster
|
||||
}
|
||||
|
||||
// UnregisterProxyFromCluster indicates an expected call of UnregisterProxyFromCluster.
|
||||
func (mr *MockControllerMockRecorder) UnregisterProxyFromCluster(ctx, clusterAddr, proxyID interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) UnregisterProxyFromCluster(ctx, clusterAddr, proxyID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnregisterProxyFromCluster", reflect.TypeOf((*MockController)(nil).UnregisterProxyFromCluster), ctx, clusterAddr, proxyID)
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package service
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./interface.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package service is a generated GoMock package.
|
||||
package service
|
||||
@@ -8,14 +13,15 @@ import (
|
||||
context "context"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
proxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -45,7 +51,7 @@ func (m *MockManager) CreateService(ctx context.Context, accountID, userID strin
|
||||
}
|
||||
|
||||
// CreateService indicates an expected call of CreateService.
|
||||
func (mr *MockManagerMockRecorder) CreateService(ctx, accountID, userID, service interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CreateService(ctx, accountID, userID, service any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateService", reflect.TypeOf((*MockManager)(nil).CreateService), ctx, accountID, userID, service)
|
||||
}
|
||||
@@ -60,7 +66,7 @@ func (m *MockManager) CreateServiceFromPeer(ctx context.Context, accountID, peer
|
||||
}
|
||||
|
||||
// CreateServiceFromPeer indicates an expected call of CreateServiceFromPeer.
|
||||
func (mr *MockManagerMockRecorder) CreateServiceFromPeer(ctx, accountID, peerID, req interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CreateServiceFromPeer(ctx, accountID, peerID, req any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateServiceFromPeer", reflect.TypeOf((*MockManager)(nil).CreateServiceFromPeer), ctx, accountID, peerID, req)
|
||||
}
|
||||
@@ -74,7 +80,7 @@ func (m *MockManager) DeleteAccountCluster(ctx context.Context, accountID, userI
|
||||
}
|
||||
|
||||
// DeleteAccountCluster indicates an expected call of DeleteAccountCluster.
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, accountID, userID, clusterAddress interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, accountID, userID, clusterAddress any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, accountID, userID, clusterAddress)
|
||||
}
|
||||
@@ -88,7 +94,7 @@ func (m *MockManager) DeleteAllServices(ctx context.Context, accountID, userID s
|
||||
}
|
||||
|
||||
// DeleteAllServices indicates an expected call of DeleteAllServices.
|
||||
func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAllServices", reflect.TypeOf((*MockManager)(nil).DeleteAllServices), ctx, accountID, userID)
|
||||
}
|
||||
@@ -102,7 +108,7 @@ func (m *MockManager) DeleteService(ctx context.Context, accountID, userID, serv
|
||||
}
|
||||
|
||||
// DeleteService indicates an expected call of DeleteService.
|
||||
func (mr *MockManagerMockRecorder) DeleteService(ctx, accountID, userID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteService(ctx, accountID, userID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteService", reflect.TypeOf((*MockManager)(nil).DeleteService), ctx, accountID, userID, serviceID)
|
||||
}
|
||||
@@ -117,7 +123,7 @@ func (m *MockManager) GetAccountServices(ctx context.Context, accountID string)
|
||||
}
|
||||
|
||||
// GetAccountServices indicates an expected call of GetAccountServices.
|
||||
func (mr *MockManagerMockRecorder) GetAccountServices(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAccountServices(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountServices", reflect.TypeOf((*MockManager)(nil).GetAccountServices), ctx, accountID)
|
||||
}
|
||||
@@ -132,7 +138,7 @@ func (m *MockManager) GetAllServices(ctx context.Context, accountID, userID stri
|
||||
}
|
||||
|
||||
// GetAllServices indicates an expected call of GetAllServices.
|
||||
func (mr *MockManagerMockRecorder) GetAllServices(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAllServices(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllServices", reflect.TypeOf((*MockManager)(nil).GetAllServices), ctx, accountID, userID)
|
||||
}
|
||||
@@ -147,7 +153,7 @@ func (m *MockManager) GetClusters(ctx context.Context, accountID, userID string)
|
||||
}
|
||||
|
||||
// GetClusters indicates an expected call of GetClusters.
|
||||
func (mr *MockManagerMockRecorder) GetClusters(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetClusters(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusters", reflect.TypeOf((*MockManager)(nil).GetClusters), ctx, accountID, userID)
|
||||
}
|
||||
@@ -162,7 +168,7 @@ func (m *MockManager) GetGlobalServices(ctx context.Context) ([]*Service, error)
|
||||
}
|
||||
|
||||
// GetGlobalServices indicates an expected call of GetGlobalServices.
|
||||
func (mr *MockManagerMockRecorder) GetGlobalServices(ctx interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetGlobalServices(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetGlobalServices", reflect.TypeOf((*MockManager)(nil).GetGlobalServices), ctx)
|
||||
}
|
||||
@@ -177,7 +183,7 @@ func (m *MockManager) GetService(ctx context.Context, accountID, userID, service
|
||||
}
|
||||
|
||||
// GetService indicates an expected call of GetService.
|
||||
func (mr *MockManagerMockRecorder) GetService(ctx, accountID, userID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetService(ctx, accountID, userID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetService", reflect.TypeOf((*MockManager)(nil).GetService), ctx, accountID, userID, serviceID)
|
||||
}
|
||||
@@ -192,7 +198,7 @@ func (m *MockManager) GetServiceByDomain(ctx context.Context, domain string) (*S
|
||||
}
|
||||
|
||||
// GetServiceByDomain indicates an expected call of GetServiceByDomain.
|
||||
func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByDomain", reflect.TypeOf((*MockManager)(nil).GetServiceByDomain), ctx, domain)
|
||||
}
|
||||
@@ -207,7 +213,7 @@ func (m *MockManager) GetServiceByID(ctx context.Context, accountID, serviceID s
|
||||
}
|
||||
|
||||
// GetServiceByID indicates an expected call of GetServiceByID.
|
||||
func (mr *MockManagerMockRecorder) GetServiceByID(ctx, accountID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetServiceByID(ctx, accountID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByID", reflect.TypeOf((*MockManager)(nil).GetServiceByID), ctx, accountID, serviceID)
|
||||
}
|
||||
@@ -222,7 +228,7 @@ func (m *MockManager) GetServiceIDByTargetID(ctx context.Context, accountID, res
|
||||
}
|
||||
|
||||
// GetServiceIDByTargetID indicates an expected call of GetServiceIDByTargetID.
|
||||
func (mr *MockManagerMockRecorder) GetServiceIDByTargetID(ctx, accountID, resourceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetServiceIDByTargetID(ctx, accountID, resourceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceIDByTargetID", reflect.TypeOf((*MockManager)(nil).GetServiceIDByTargetID), ctx, accountID, resourceID)
|
||||
}
|
||||
@@ -236,7 +242,7 @@ func (m *MockManager) ReloadAllServicesForAccount(ctx context.Context, accountID
|
||||
}
|
||||
|
||||
// ReloadAllServicesForAccount indicates an expected call of ReloadAllServicesForAccount.
|
||||
func (mr *MockManagerMockRecorder) ReloadAllServicesForAccount(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ReloadAllServicesForAccount(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReloadAllServicesForAccount", reflect.TypeOf((*MockManager)(nil).ReloadAllServicesForAccount), ctx, accountID)
|
||||
}
|
||||
@@ -250,7 +256,7 @@ func (m *MockManager) ReloadService(ctx context.Context, accountID, serviceID st
|
||||
}
|
||||
|
||||
// ReloadService indicates an expected call of ReloadService.
|
||||
func (mr *MockManagerMockRecorder) ReloadService(ctx, accountID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ReloadService(ctx, accountID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReloadService", reflect.TypeOf((*MockManager)(nil).ReloadService), ctx, accountID, serviceID)
|
||||
}
|
||||
@@ -264,7 +270,7 @@ func (m *MockManager) RenewServiceFromPeer(ctx context.Context, accountID, peerI
|
||||
}
|
||||
|
||||
// RenewServiceFromPeer indicates an expected call of RenewServiceFromPeer.
|
||||
func (mr *MockManagerMockRecorder) RenewServiceFromPeer(ctx, accountID, peerID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) RenewServiceFromPeer(ctx, accountID, peerID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RenewServiceFromPeer", reflect.TypeOf((*MockManager)(nil).RenewServiceFromPeer), ctx, accountID, peerID, serviceID)
|
||||
}
|
||||
@@ -278,7 +284,7 @@ func (m *MockManager) SetCertificateIssuedAt(ctx context.Context, accountID, ser
|
||||
}
|
||||
|
||||
// SetCertificateIssuedAt indicates an expected call of SetCertificateIssuedAt.
|
||||
func (mr *MockManagerMockRecorder) SetCertificateIssuedAt(ctx, accountID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetCertificateIssuedAt(ctx, accountID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetCertificateIssuedAt", reflect.TypeOf((*MockManager)(nil).SetCertificateIssuedAt), ctx, accountID, serviceID)
|
||||
}
|
||||
@@ -292,7 +298,7 @@ func (m *MockManager) SetStatus(ctx context.Context, accountID, serviceID string
|
||||
}
|
||||
|
||||
// SetStatus indicates an expected call of SetStatus.
|
||||
func (mr *MockManagerMockRecorder) SetStatus(ctx, accountID, serviceID, status interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetStatus(ctx, accountID, serviceID, status any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetStatus", reflect.TypeOf((*MockManager)(nil).SetStatus), ctx, accountID, serviceID, status)
|
||||
}
|
||||
@@ -304,7 +310,7 @@ func (m *MockManager) StartExposeReaper(ctx context.Context) {
|
||||
}
|
||||
|
||||
// StartExposeReaper indicates an expected call of StartExposeReaper.
|
||||
func (mr *MockManagerMockRecorder) StartExposeReaper(ctx interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) StartExposeReaper(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StartExposeReaper", reflect.TypeOf((*MockManager)(nil).StartExposeReaper), ctx)
|
||||
}
|
||||
@@ -318,7 +324,7 @@ func (m *MockManager) StopServiceFromPeer(ctx context.Context, accountID, peerID
|
||||
}
|
||||
|
||||
// StopServiceFromPeer indicates an expected call of StopServiceFromPeer.
|
||||
func (mr *MockManagerMockRecorder) StopServiceFromPeer(ctx, accountID, peerID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) StopServiceFromPeer(ctx, accountID, peerID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StopServiceFromPeer", reflect.TypeOf((*MockManager)(nil).StopServiceFromPeer), ctx, accountID, peerID, serviceID)
|
||||
}
|
||||
@@ -333,7 +339,7 @@ func (m *MockManager) UpdateService(ctx context.Context, accountID, userID strin
|
||||
}
|
||||
|
||||
// UpdateService indicates an expected call of UpdateService.
|
||||
func (mr *MockManagerMockRecorder) UpdateService(ctx, accountID, userID, service interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) UpdateService(ctx, accountID, userID, service any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateService", reflect.TypeOf((*MockManager)(nil).UpdateService), ctx, accountID, userID, service)
|
||||
}
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"time"
|
||||
|
||||
cachestore "github.com/eko/gocache/lib/v4/store"
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user