Compare commits

..

1 Commits

Author SHA1 Message Date
pascal
7ed3737cda revert component types 2026-07-24 15:16:55 +02:00
92 changed files with 1270 additions and 6211 deletions

View File

@@ -5,13 +5,6 @@ on:
schedule:
- cron: "0 3 * * *"
workflow_dispatch:
inputs:
bedrock_model:
description: >-
Bedrock inference-profile id to drive the matrix with, exactly as
AWS issues it. Leave empty for the Sonnet 4.6 default.
required: false
default: ""
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
@@ -69,8 +62,6 @@ jobs:
CLOUDFLARE_TOKEN: ${{ secrets.E2E_CLOUDFLARE_TOKEN }}
AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.E2E_AWS_BEARER_TOKEN_BEDROCK }}
AWS_REGION: ${{ secrets.E2E_AWS_REGION }}
# Bedrock model override: dispatch input wins, then the repo variable, else the test default.
AWS_BEDROCK_MODEL: ${{ inputs.bedrock_model || vars.E2E_AWS_BEDROCK_MODEL }}
# Vertex (Anthropic-on-Vertex): SA + project required; region defaults
# to "global", model to a pinned claude snapshot.
GOOGLE_VERTEX_SA_BASE64: ${{ secrets.E2E_GOOGLE_VERTEX_SA_BASE64 }}

View File

@@ -464,8 +464,6 @@ func Test_RemovePeer(t *testing.T) {
}
func Test_ConnectPeers(t *testing.T) {
t.Setenv("NB_DISABLE_EBPF_WG_PROXY", "true")
peer1ifaceName := fmt.Sprintf("utun%d", WgIntNumber+400)
peer1wgIP := netip.MustParsePrefix("10.99.99.17/30")
peer1Key, _ := wgtypes.GeneratePrivateKey()
@@ -507,8 +505,12 @@ func Test_ConnectPeers(t *testing.T) {
t.Fatal(err)
}
localIP1 := "127.0.0.1"
peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP1, peer1wgPort))
localIP, err := getLocalIP()
if err != nil {
t.Fatal(err)
}
peer1endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer1wgPort))
if err != nil {
t.Fatal(err)
}
@@ -544,8 +546,7 @@ func Test_ConnectPeers(t *testing.T) {
t.Fatal(err)
}
localIP2 := "127.0.0.1"
peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP2, peer2wgPort))
peer2endpoint, err := net.ResolveUDPAddr("udp", fmt.Sprintf("%s:%d", localIP, peer2wgPort))
if err != nil {
t.Fatal(err)
}
@@ -620,3 +621,28 @@ func getPeer(ifaceName, peerPubKey string) (wgtypes.Peer, error) {
}
return wgtypes.Peer{}, fmt.Errorf("peer not found")
}
func getLocalIP() (string, error) {
// Get all interfaces
addrs, err := net.InterfaceAddrs()
if err != nil {
return "", err
}
for _, addr := range addrs {
ipNet, ok := addr.(*net.IPNet)
if !ok {
continue
}
if ipNet.IP.IsLoopback() {
continue
}
if ipNet.IP.To4() == nil {
continue
}
return ipNet.IP.String(), nil
}
return "", fmt.Errorf("no local IP found")
}

View File

@@ -9,220 +9,12 @@ import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"github.com/netbirdio/netbird/e2e/harness"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// per1k is a model's published USD rates per 1k tokens. read is the prompt-cache read rate
// (OpenAI: the cached-input discount rate); write is the cache-creation rate where one exists.
type per1k struct{ in, out, read, write float64 }
// publishedPer1k hardcodes the vendors' PUBLISHED rates for the models the live matrix can drive,
// keyed by the normalized model id the proxy stamps. Deliberately independent of the proxy's
// pricing table so a wrong embedded rate or a broken normalization fails the run.
var publishedPer1k = map[string]per1k{
"gpt-4o-mini": {0.00015, 0.0006, 0.000075, 0},
"gpt-4o": {0.0025, 0.01, 0.00125, 0},
"claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125},
"claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375},
"claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375},
"kimi-k3": {0.003, 0.015, 0.0003, 0.003}, // no published write rate: bills at the input rate
"anthropic.claude-haiku-4-5": {0.001, 0.005, 0.0001, 0.00125},
"anthropic.claude-sonnet-4-5": {0.003, 0.015, 0.0003, 0.00375},
"anthropic.claude-sonnet-4-6": {0.003, 0.015, 0.0003, 0.00375},
}
// rawCostVerificationSQL is the operator-facing double-check, run straight against the management
// sqlite store: recompute each usage row's expected total and cache cost from its own persisted
// token buckets and hardcoded published rates. OpenAI counts cached tokens as a subset of input;
// Anthropic-shape providers count cache buckets additively.
const rawCostVerificationSQL = `
WITH rates(model, in_rate, out_rate, read_rate, write_rate) AS (
VALUES
('gpt-4o-mini', 0.00015, 0.0006, 0.000075, 0.0),
('gpt-4o', 0.0025, 0.01, 0.00125, 0.0),
('claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125),
('claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375),
('claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375),
('kimi-k3', 0.003, 0.015, 0.0003, 0.003),
('anthropic.claude-haiku-4-5', 0.001, 0.005, 0.0001, 0.00125),
('anthropic.claude-sonnet-4-5', 0.003, 0.015, 0.0003, 0.00375),
('anthropic.claude-sonnet-4-6', 0.003, 0.015, 0.0003, 0.00375)
)
SELECT
u.provider,
u.model,
u.input_tokens,
u.output_tokens,
u.cached_input_tokens,
u.cache_creation_tokens,
u.input_cost_usd,
u.cached_input_cost_usd,
u.cache_creation_cost_usd,
u.output_cost_usd,
-- No cost_usd / cache_cost_usd columns are stored: both are derived from the
-- four per-bucket columns above, exactly as the API renders them.
(u.input_cost_usd + u.cached_input_cost_usd + u.cache_creation_cost_usd + u.output_cost_usd) AS cost_usd,
(u.cached_input_cost_usd + u.cache_creation_cost_usd) AS cache_cost_usd,
CASE WHEN u.provider = 'openai' THEN
(u.input_tokens - MIN(u.cached_input_tokens, u.input_tokens))*r.in_rate/1000.0
ELSE
u.input_tokens*r.in_rate/1000.0
END AS expected_input,
CASE WHEN u.provider = 'openai' THEN
MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
ELSE
u.cached_input_tokens*r.read_rate/1000.0
END AS expected_cached_input,
CASE WHEN u.provider = 'openai' THEN
0.0
ELSE
u.cache_creation_tokens*r.write_rate/1000.0
END AS expected_cache_creation,
u.output_tokens*r.out_rate/1000.0 AS expected_output,
CASE WHEN u.provider = 'openai' THEN
(u.input_tokens - MIN(u.cached_input_tokens, u.input_tokens))*r.in_rate/1000.0
+ MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
+ u.output_tokens*r.out_rate/1000.0
ELSE
u.input_tokens*r.in_rate/1000.0 + u.cached_input_tokens*r.read_rate/1000.0
+ u.cache_creation_tokens*r.write_rate/1000.0 + u.output_tokens*r.out_rate/1000.0
END AS expected_total,
CASE WHEN u.provider = 'openai' THEN
MIN(u.cached_input_tokens, u.input_tokens)*r.read_rate/1000.0
ELSE
u.cached_input_tokens*r.read_rate/1000.0 + u.cache_creation_tokens*r.write_rate/1000.0
END AS expected_cache
FROM agent_network_request_usage u
JOIN rates r ON r.model = u.model
ORDER BY u.timestamp`
// verifyUsageRowsSQL re-checks every persisted usage row directly in the management sqlite store,
// bypassing the API path — the same audit an operator can run on a production store.db.
func verifyUsageRowsSQL(t *testing.T, srv *harness.Combined) {
t.Helper()
dbPath, err := srv.SnapshotStoreDB(t.TempDir())
require.NoError(t, err, "snapshot management sqlite store")
db, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{})
require.NoError(t, err, "open store snapshot")
sqlDB, err := db.DB()
require.NoError(t, err)
defer func() { _ = sqlDB.Close() }()
rows, err := db.Raw(rawCostVerificationSQL).Rows()
require.NoError(t, err, "run raw cost verification query")
defer func() { _ = rows.Close() }()
verified := 0
for rows.Next() {
var provider, model string
var inTok, outTok, readTok, writeTok int64
var inCost, cachedInCost, cacheCreateCost, outCost, cost, cacheCost float64
var wantInput, wantCachedInput, wantCacheCreation, wantOutput, wantTotal, wantCache float64
require.NoError(t, rows.Scan(&provider, &model, &inTok, &outTok, &readTok, &writeTok,
&inCost, &cachedInCost, &cacheCreateCost, &outCost, &cost, &cacheCost,
&wantInput, &wantCachedInput, &wantCacheCreation, &wantOutput, &wantTotal, &wantCache), "scan usage row")
t.Logf("[sql] %s/%s: in=%d out=%d cache_read=%d cache_write=%d stored in/cached/create/out=$%.6f/$%.6f/$%.6f/$%.6f total=$%.6f cache=$%.6f expected total=$%.6f cache=$%.6f",
provider, model, inTok, outTok, readTok, writeTok,
inCost, cachedInCost, cacheCreateCost, outCost, cost, cacheCost, wantTotal, wantCache)
assert.InDeltaf(t, wantInput, inCost, 1e-6, "stored input_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCachedInput, cachedInCost, 1e-6, "stored cached_input_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCacheCreation, cacheCreateCost, 1e-6, "stored cache_creation_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantOutput, outCost, 1e-6, "stored output_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantTotal, cost, 1e-6, "derived cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, wantCache, cacheCost, 1e-6, "derived cache_cost_usd for %s/%s must match the published-rate recompute", provider, model)
assert.InDeltaf(t, inCost+cachedInCost+cacheCreateCost+outCost, cost, 1e-9,
"stored buckets must sum to the derived cost_usd for %s/%s", provider, model)
verified++
}
require.NoError(t, rows.Err(), "iterate usage rows")
require.Positive(t, verified, "raw SQL check must cover at least one usage row")
t.Logf("[sql] verified %d usage rows in store.db against published rates", verified)
gwRows, err := db.Raw(`SELECT model,
(input_cost_usd + cached_input_cost_usd + cache_creation_cost_usd + output_cost_usd) AS cost_usd
FROM agent_network_request_usage WHERE model LIKE '%/%'`).Rows()
require.NoError(t, err, "query gateway-prefixed usage rows")
defer func() { _ = gwRows.Close() }()
for gwRows.Next() {
var model string
var cost float64
require.NoError(t, gwRows.Scan(&model, &cost), "scan gateway usage row")
t.Logf("[sql] gateway %s: stored=$%.6f (must be 0 — deliberately unpriced)", model, cost)
assert.Zerof(t, cost, "gateway-prefixed model %q must store cost 0, never a guessed rate", model)
}
require.NoError(t, gwRows.Err(), "iterate gateway usage rows")
}
// validateAccessLogCost recomputes a live access-log row's expected total and cache cost from the
// published per-1k rates and the row's persisted token buckets, and asserts both stored values.
// Gateway-prefixed model ids the proxy deliberately does not price must store cost 0.
func validateAccessLogCost(t *testing.T, pc providerCase, row api.AgentNetworkAccessLog) {
t.Helper()
model := catalogModel(pc)
provider := ""
if row.Provider != nil {
provider = *row.Provider
}
t.Logf("[cost] %s: provider=%s model=%s in=%d out=%d total=%d cache_read=%d cache_write=%d cost=$%.6f cache_cost=$%.6f",
pc.name, provider, model, row.InputTokens, row.OutputTokens, row.TotalTokens,
row.CachedInputTokens, row.CacheCreationTokens, row.CostUsd, row.CacheCostUsd)
rates, known := publishedPer1k[model]
if !known {
if strings.Contains(model, "/") {
assert.Zerof(t, row.CostUsd, "gateway-prefixed model %q is not priced so the cost meter must skip (cost 0)", model)
return
}
t.Logf("[cost] %s: no published rate on file for model %q (env-overridden?); skipping cost validation", pc.name, model)
return
}
// input_tokens may legitimately be 0: Moonshot/Kimi reports fully cached prompts under the cache
// buckets only. Output and total must always be present on a priced row.
require.Positive(t, row.OutputTokens, "priced row must carry output tokens")
require.Positive(t, row.TotalTokens, "priced row must carry total tokens")
var wantInput, wantCachedInput, wantCacheCreation float64
if provider == "openai" {
cached := min(row.CachedInputTokens, row.InputTokens) // cached is a subset of input
wantInput = float64(row.InputTokens-cached) / 1000 * rates.in
wantCachedInput = float64(cached) / 1000 * rates.read
// OpenAI has no cache-write bucket; wantCacheCreation stays 0.
} else {
// Anthropic / Bedrock shape: cache buckets are additive to input_tokens.
wantInput = float64(row.InputTokens) / 1000 * rates.in
wantCachedInput = float64(row.CachedInputTokens) / 1000 * rates.read
wantCacheCreation = float64(row.CacheCreationTokens) / 1000 * rates.write
}
wantOutput := float64(row.OutputTokens) / 1000 * rates.out
wantCache := wantCachedInput + wantCacheCreation
wantTotal := wantInput + wantCache + wantOutput
t.Logf("[cost] %s: expecting input=$%.6f cached_input=$%.6f cache_creation=$%.6f output=$%.6f total=$%.6f cache=$%.6f from published rates",
pc.name, wantInput, wantCachedInput, wantCacheCreation, wantOutput, wantTotal, wantCache)
assert.InDeltaf(t, wantInput, row.InputCostUsd, 1e-6, "stored input_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCachedInput, row.CachedInputCostUsd, 1e-6, "stored cached_input_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCacheCreation, row.CacheCreationCostUsd, 1e-6, "stored cache_creation_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantOutput, row.OutputCostUsd, 1e-6, "stored output_cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantTotal, row.CostUsd, 1e-6, "derived cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, wantCache, row.CacheCostUsd, 1e-6, "derived cache_cost_usd for %s (%s)", pc.name, model)
// The aggregates must be exactly the sum of the stored components, not an
// independently-computed figure that could drift from the breakdown.
assert.InDeltaf(t, row.InputCostUsd+row.CachedInputCostUsd+row.CacheCreationCostUsd+row.OutputCostUsd,
row.CostUsd, 1e-9, "stored buckets must sum to the derived cost_usd for %s (%s)", pc.name, model)
assert.InDeltaf(t, row.CachedInputCostUsd+row.CacheCreationCostUsd,
row.CacheCostUsd, 1e-9, "stored cache buckets must sum to the derived cache_cost_usd for %s (%s)", pc.name, model)
}
// providerCase is one entry in the live provider matrix. The same scenario runs
// for every available provider; availability is keyed off env vars so the suite
// covers whatever credentials are present (source ~/.llm-keys locally / set the
@@ -324,12 +116,12 @@ func availableProviders() []providerCase {
if region == "" {
region = "eu-central-1"
}
// A valid Bedrock inference-profile id, overridable per account (AWS_BEDROCK_MODEL, also the
// workflow's bedrock_model dispatch input). `global.` profiles work from any region. Defaults to
// Sonnet 4.6, whose id convention dropped the -YYYYMMDD-v1:0 suffix that Haiku 4.5 still carries.
// A valid Bedrock inference-profile id (region prefix + date + version),
// overridable per account. `global.` profiles can be invoked from any
// region; set AWS_BEDROCK_MODEL to match the enabled profile for the token.
model := os.Getenv("AWS_BEDROCK_MODEL")
if model == "" {
model = "global.anthropic.claude-sonnet-4-6"
model = "global.anthropic.claude-haiku-4-5-20251001-v1:0"
}
ps = append(ps, providerCase{name: "bedrock", catalogID: "bedrock_api", upstream: "https://bedrock-runtime." + region + ".amazonaws.com", apiKey: k, model: model, kind: harness.WireBedrock})
}
@@ -465,10 +257,6 @@ func TestProvidersMatrix(t *testing.T) {
// session id and confirm the marker propagated end-to-end.
sessionID := "e2e-session-" + pc.name
// A long-form prompt so completions carry realistic token counts for cost validation;
// max_tokens in the harness bodies (2048) lets the full answer through.
const matrixPrompt = "explain GitHub workflow in 1000 words"
// Retry briefly to absorb tunnel/DNS jitter on the first call.
var code int
var body string
@@ -479,11 +267,11 @@ func TestProvidersMatrix(t *testing.T) {
var cerr error
switch pc.kind {
case harness.WireVertex:
c, b, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, pc.project, pc.region, pc.model, matrixPrompt, sessionID)
c, b, cerr = cl.Vertex(ctx, settings.Endpoint, proxyIP, pc.project, pc.region, pc.model, "Reply with exactly: pong", sessionID)
case harness.WireBedrock:
c, b, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, pc.model, matrixPrompt, sessionID)
c, b, cerr = cl.Bedrock(ctx, settings.Endpoint, proxyIP, pc.model, "Reply with exactly: pong", sessionID)
default:
c, b, cerr = cl.ChatPrefixed(ctx, settings.Endpoint, proxyIP, pc.pathPrefix, pc.kind, pc.model, matrixPrompt, sessionID)
c, b, cerr = cl.ChatPrefixed(ctx, settings.Endpoint, proxyIP, pc.pathPrefix, pc.kind, pc.model, "Reply with exactly: pong", sessionID)
}
if cerr == nil {
code, body = c, b
@@ -502,7 +290,6 @@ func TestProvidersMatrix(t *testing.T) {
// The session id sent as x-session-id must round-trip into the
// access-log row for this provider.
var row api.AgentNetworkAccessLog
require.Eventually(t, func() bool {
logs, lerr := srv.ListAccessLogs(ctx)
if lerr != nil {
@@ -510,15 +297,11 @@ func TestProvidersMatrix(t *testing.T) {
}
for _, r := range logs.Data {
if r.SessionId != nil && *r.SessionId == sessionID {
row = r
return true
}
}
return false
}, 30*time.Second, 2*time.Second, "session id %q must be recorded in an access-log row for %s", sessionID, pc.name)
// Stored total and cache cost must match the published rates applied to the row's buckets.
validateAccessLogCost(t, pc, row)
})
}
@@ -539,7 +322,4 @@ func TestProvidersMatrix(t *testing.T) {
}
return false
}, 60*time.Second, 3*time.Second, "consumption must be recorded with positive token counts after live traffic")
// Final raw-SQL audit: bypass the API and re-verify every persisted usage row in the store.
verifyUsageRowsSQL(t, srv)
}

View File

@@ -256,7 +256,7 @@ func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefi
case WireMessages:
path = "/v1/messages"
headers = []string{"anthropic-version: 2023-06-01"}
body = fmt.Sprintf(`{"model":%q,"max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, model, prompt)
body = fmt.Sprintf(`{"model":%q,"max_tokens":64,"messages":[{"role":"user","content":%q}]}`, model, prompt)
default:
path = "/v1/chat/completions"
body = fmt.Sprintf(`{"model":%q,"messages":[{"role":"user","content":%q}]}`, model, prompt)
@@ -271,7 +271,7 @@ func (cl *Client) ChatPrefixed(ctx context.Context, endpoint, proxyIP, pathPrefi
// is sent as the universal x-session-id header the proxy records.
func (cl *Client) Vertex(ctx context.Context, endpoint, proxyIP, project, region, model, prompt, sessionID string) (int, string, error) {
path := fmt.Sprintf("/v1/projects/%s/locations/%s/publishers/anthropic/models/%s:rawPredict", project, region, model)
body := fmt.Sprintf(`{"anthropic_version":"vertex-2023-10-16","max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, prompt)
body := fmt.Sprintf(`{"anthropic_version":"vertex-2023-10-16","max_tokens":64,"messages":[{"role":"user","content":%q}]}`, prompt)
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(nil, sessionID))
}
@@ -282,7 +282,7 @@ func (cl *Client) Vertex(ctx context.Context, endpoint, proxyIP, project, region
// header the proxy records.
func (cl *Client) Bedrock(ctx context.Context, endpoint, proxyIP, model, prompt, sessionID string) (int, string, error) {
path := "/model/" + model + "/invoke"
body := fmt.Sprintf(`{"anthropic_version":"bedrock-2023-05-31","max_tokens":2048,"messages":[{"role":"user","content":%q}]}`, prompt)
body := fmt.Sprintf(`{"anthropic_version":"bedrock-2023-05-31","max_tokens":64,"messages":[{"role":"user","content":%q}]}`, prompt)
return cl.post(ctx, endpoint, proxyIP, path, body, withSessionID(nil, sessionID))
}
@@ -341,7 +341,7 @@ func (cl *Client) Terminate(ctx context.Context) error {
return cl.container.Terminate(ctx)
}
// containerLogs reads up to 4 MiB of a container's logs for diagnostics — enough for a whole provider-matrix run.
// containerLogs reads up to 256 KiB of a container's logs for diagnostics.
func containerLogs(ctx context.Context, c testcontainers.Container) string {
if c == nil {
return ""
@@ -351,6 +351,6 @@ func containerLogs(ctx context.Context, c testcontainers.Container) string {
return fmt.Sprintf("<logs error: %v>", err)
}
defer r.Close()
b, _ := io.ReadAll(io.LimitReader(r, 4<<20))
b, _ := io.ReadAll(io.LimitReader(r, 256<<10))
return string(b)
}

View File

@@ -221,29 +221,6 @@ func (c *Combined) CreateProxyTokenCLI(ctx context.Context, name string) (string
return "", fmt.Errorf("token not found in CLI output: %s", string(out))
}
// SnapshotStoreDB copies the management sqlite store (with WAL/SHM sidecars) out of the bind-mounted
// data dir into dstDir and returns the copy's path; reading a copy avoids locking against live writes.
func (c *Combined) SnapshotStoreDB(dstDir string) (string, error) {
src := filepath.Join(c.workDir, "data", "store.db")
if _, err := os.Stat(src); err != nil {
return "", fmt.Errorf("management store not found at %s: %w", src, err)
}
dst := filepath.Join(dstDir, "store.db")
for _, suffix := range []string{"", "-wal", "-shm"} {
data, err := os.ReadFile(src + suffix)
if err != nil {
if os.IsNotExist(err) && suffix != "" {
continue // sidecar only exists in WAL mode
}
return "", fmt.Errorf("read %s: %w", src+suffix, err)
}
if err := os.WriteFile(dst+suffix, data, 0o600); err != nil {
return "", fmt.Errorf("write %s: %w", dst+suffix, err)
}
}
return dst, nil
}
// Logs returns the combined server container logs, for diagnostics.
func (c *Combined) Logs(ctx context.Context) string {
return containerLogs(ctx, c.container)

View File

@@ -18,26 +18,21 @@ import (
// contract between the proxy and management; management flattens them into
// queryable columns. Keep in sync with the proxy side.
const (
metaKeyProvider = "llm.provider"
metaKeyModel = "llm.model"
metaKeyResolvedProviderID = "llm.resolved_provider_id"
metaKeySelectedPolicyID = "llm.selected_policy_id"
metaKeyPolicyDecision = "llm_policy.decision"
metaKeyPolicyReason = "llm_policy.reason"
metaKeyInputTokens = "llm.input_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyOutputTokens = "llm.output_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyTotalTokens = "llm.total_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCachedInputTokens = "llm.cached_input_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCacheCreationTokens = "llm.cache_creation_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCostUSDInput = "cost.usd_input"
metaKeyCostUSDCachedInput = "cost.usd_cached_input"
metaKeyCostUSDCacheCreate = "cost.usd_cache_creation"
metaKeyCostUSDOutput = "cost.usd_output"
metaKeyStream = "llm.stream"
metaKeySessionID = "llm.session_id"
metaKeyAuthorisingGroups = "llm.authorising_groups"
metaKeyRequestPrompt = "llm.request_prompt"
metaKeyResponseCompletion = "llm.response_completion"
metaKeyProvider = "llm.provider"
metaKeyModel = "llm.model"
metaKeyResolvedProviderID = "llm.resolved_provider_id"
metaKeySelectedPolicyID = "llm.selected_policy_id"
metaKeyPolicyDecision = "llm_policy.decision"
metaKeyPolicyReason = "llm_policy.reason"
metaKeyInputTokens = "llm.input_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyOutputTokens = "llm.output_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyTotalTokens = "llm.total_tokens" //nolint:gosec // metadata key name, not a credential
metaKeyCostUSDTotal = "cost.usd_total"
metaKeyStream = "llm.stream"
metaKeySessionID = "llm.session_id"
metaKeyAuthorisingGroups = "llm.authorising_groups"
metaKeyRequestPrompt = "llm.request_prompt"
metaKeyResponseCompletion = "llm.response_completion"
)
// IngestAccessLog flattens the metadata-bearing reverse-proxy access-log entry
@@ -113,25 +108,20 @@ func flattenAccessLog(e *accesslogs.AccessLogEntry) (*types.AgentNetworkAccessLo
BytesUpload: e.BytesUpload,
BytesDownload: e.BytesDownload,
Provider: meta[metaKeyProvider],
Model: meta[metaKeyModel],
SessionID: meta[metaKeySessionID],
ResolvedProviderID: meta[metaKeyResolvedProviderID],
SelectedPolicyID: meta[metaKeySelectedPolicyID],
Decision: meta[metaKeyPolicyDecision],
DenyReason: meta[metaKeyPolicyReason],
InputTokens: parseMetaInt(meta, metaKeyInputTokens),
OutputTokens: parseMetaInt(meta, metaKeyOutputTokens),
TotalTokens: parseMetaInt(meta, metaKeyTotalTokens),
CachedInputTokens: parseMetaInt(meta, metaKeyCachedInputTokens),
CacheCreationTokens: parseMetaInt(meta, metaKeyCacheCreationTokens),
InputCostUSD: parseMetaFloat(meta, metaKeyCostUSDInput),
CachedInputCostUSD: parseMetaFloat(meta, metaKeyCostUSDCachedInput),
CacheCreationCostUSD: parseMetaFloat(meta, metaKeyCostUSDCacheCreate),
OutputCostUSD: parseMetaFloat(meta, metaKeyCostUSDOutput),
Stream: parseMetaBool(meta, metaKeyStream),
RequestPrompt: meta[metaKeyRequestPrompt],
ResponseCompletion: meta[metaKeyResponseCompletion],
Provider: meta[metaKeyProvider],
Model: meta[metaKeyModel],
SessionID: meta[metaKeySessionID],
ResolvedProviderID: meta[metaKeyResolvedProviderID],
SelectedPolicyID: meta[metaKeySelectedPolicyID],
Decision: meta[metaKeyPolicyDecision],
DenyReason: meta[metaKeyPolicyReason],
InputTokens: parseMetaInt(meta, metaKeyInputTokens),
OutputTokens: parseMetaInt(meta, metaKeyOutputTokens),
TotalTokens: parseMetaInt(meta, metaKeyTotalTokens),
CostUSD: parseMetaFloat(meta, metaKeyCostUSDTotal),
Stream: parseMetaBool(meta, metaKeyStream),
RequestPrompt: meta[metaKeyRequestPrompt],
ResponseCompletion: meta[metaKeyResponseCompletion],
}
var groups []types.AgentNetworkAccessLogGroup
@@ -150,23 +140,18 @@ func flattenAccessLog(e *accesslogs.AccessLogEntry) (*types.AgentNetworkAccessLo
// log's ID so the two correlate.
func usageFromFlattenedLog(e *types.AgentNetworkAccessLog, groups []types.AgentNetworkAccessLogGroup) (*types.AgentNetworkUsage, []types.AgentNetworkUsageGroup) {
usage := &types.AgentNetworkUsage{
ID: e.ID,
AccountID: e.AccountID,
Timestamp: e.Timestamp,
UserID: e.UserID,
ResolvedProviderID: e.ResolvedProviderID,
Provider: e.Provider,
Model: e.Model,
SessionID: e.SessionID,
InputTokens: e.InputTokens,
OutputTokens: e.OutputTokens,
TotalTokens: e.TotalTokens,
CachedInputTokens: e.CachedInputTokens,
CacheCreationTokens: e.CacheCreationTokens,
InputCostUSD: e.InputCostUSD,
CachedInputCostUSD: e.CachedInputCostUSD,
CacheCreationCostUSD: e.CacheCreationCostUSD,
OutputCostUSD: e.OutputCostUSD,
ID: e.ID,
AccountID: e.AccountID,
Timestamp: e.Timestamp,
UserID: e.UserID,
ResolvedProviderID: e.ResolvedProviderID,
Provider: e.Provider,
Model: e.Model,
SessionID: e.SessionID,
InputTokens: e.InputTokens,
OutputTokens: e.OutputTokens,
TotalTokens: e.TotalTokens,
CostUSD: e.CostUSD,
}
usageGroups := make([]types.AgentNetworkUsageGroup, 0, len(groups))

View File

@@ -28,22 +28,17 @@ func newIngestTestEntry() *accesslogs.AccessLogEntry {
UserId: "user-1",
AgentNetwork: true,
Metadata: map[string]string{
metaKeyProvider: "openai",
metaKeyModel: "gpt-5.4",
metaKeyResolvedProviderID: "prov-1",
metaKeySessionID: "sess-1",
metaKeyInputTokens: "100",
metaKeyOutputTokens: "50",
metaKeyTotalTokens: "1174",
metaKeyCachedInputTokens: "256",
metaKeyCacheCreationTokens: "768",
metaKeyCostUSDInput: "0.0071",
metaKeyCostUSDCachedInput: "0.0009",
metaKeyCostUSDCacheCreate: "0.0020",
metaKeyCostUSDOutput: "0.0023",
metaKeyStream: "true",
metaKeyRequestPrompt: "hello",
metaKeyResponseCompletion: "world",
metaKeyProvider: "openai",
metaKeyModel: "gpt-5.4",
metaKeyResolvedProviderID: "prov-1",
metaKeySessionID: "sess-1",
metaKeyInputTokens: "100",
metaKeyOutputTokens: "50",
metaKeyTotalTokens: "150",
metaKeyCostUSDTotal: "0.0123",
metaKeyStream: "true",
metaKeyRequestPrompt: "hello",
metaKeyResponseCompletion: "world",
// repeated id must be de-duplicated before the group rows insert.
metaKeyAuthorisingGroups: "grp-eng,grp-eng,grp-ops",
},
@@ -70,19 +65,7 @@ func TestIngestAccessLog_RealStore_LogCollectionOff(t *testing.T) {
require.Len(t, usage, 1, "usage row must be written even with log collection off")
assert.Equal(t, int64(100), usage[0].InputTokens, "input tokens must round-trip from metadata")
assert.Equal(t, int64(50), usage[0].OutputTokens, "output tokens must round-trip from metadata")
assert.Equal(t, int64(256), usage[0].CachedInputTokens, "cache-read tokens must round-trip from metadata")
assert.Equal(t, int64(768), usage[0].CacheCreationTokens, "cache-write tokens must round-trip from metadata")
// The per-bucket breakdown is the only cost state stored, and must survive
// the write/read cycle as real columns — usage rows are the only cost
// record for accounts with log collection off, so a dropped column here
// loses the split permanently.
assert.InDelta(t, 0.0071, usage[0].InputCostUSD, 1e-9, "input cost must round-trip from metadata")
assert.InDelta(t, 0.0009, usage[0].CachedInputCostUSD, 1e-9, "cache-read cost must round-trip from metadata")
assert.InDelta(t, 0.0020, usage[0].CacheCreationCostUSD, 1e-9, "cache-write cost must round-trip from metadata")
assert.InDelta(t, 0.0023, usage[0].OutputCostUSD, 1e-9, "output cost must round-trip from metadata")
// Aggregates are derived from the stored columns, never stored themselves.
assert.InDelta(t, 0.0123, usage[0].TotalCostUSD(), 1e-9, "total is derived from the stored buckets")
assert.InDelta(t, 0.0029, usage[0].CacheCostUSD(), 1e-9, "cache cost is derived from the two cache buckets")
assert.InDelta(t, 0.0123, usage[0].CostUSD, 1e-9, "cost must round-trip from metadata")
logs, _, err := s.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, testAccountID, types.AgentNetworkAccessLogFilter{})
require.NoError(t, err)
@@ -113,14 +96,6 @@ func TestIngestAccessLog_RealStore_LogCollectionOn(t *testing.T) {
require.Equal(t, int64(1), total, "exactly one access-log row expected")
require.Len(t, logs, 1, "full access-log row must be written when log collection is on")
assert.Equal(t, "gpt-5.4", logs[0].Model, "model must flatten from metadata")
assert.Equal(t, int64(256), logs[0].CachedInputTokens, "cache-read tokens must flatten from metadata")
assert.Equal(t, int64(768), logs[0].CacheCreationTokens, "cache-write tokens must flatten from metadata")
assert.InDelta(t, 0.0029, logs[0].CacheCostUSD(), 1e-9, "cache cost is derived from the two cache buckets")
assert.InDelta(t, 0.0123, logs[0].TotalCostUSD(), 1e-9, "total is derived from the stored buckets")
assert.InDelta(t, 0.0071, logs[0].InputCostUSD, 1e-9, "input cost must flatten from metadata")
assert.InDelta(t, 0.0009, logs[0].CachedInputCostUSD, 1e-9, "cache-read cost must flatten from metadata")
assert.InDelta(t, 0.0020, logs[0].CacheCreationCostUSD, 1e-9, "cache-write cost must flatten from metadata")
assert.InDelta(t, 0.0023, logs[0].OutputCostUSD, 1e-9, "output cost must flatten from metadata")
assert.Equal(t, "hello", logs[0].RequestPrompt, "prompt must be retained when log collection is on")
assert.Equal(t, "world", logs[0].ResponseCompletion, "completion must be retained when log collection is on")
assert.True(t, logs[0].Stream, "stream flag must flatten from metadata")

View File

@@ -38,7 +38,7 @@ func accessLogRow(id, sessionID string, ts time.Time, opts ...func(*types.AgentN
InputTokens: 100,
OutputTokens: 50,
TotalTokens: 150,
InputCostUSD: 0.01,
CostUSD: 0.01,
}
for _, o := range opts {
o(e)
@@ -74,7 +74,7 @@ func withTokens(in, out, total int64, cost float64) func(*types.AgentNetworkAcce
e.InputTokens = in
e.OutputTokens = out
e.TotalTokens = total
e.InputCostUSD = cost
e.CostUSD = cost
}
}
@@ -155,7 +155,7 @@ func TestAccessLogSessions_FoldAndAggregate(t *testing.T) {
assert.Equal(t, int64(310), a.InputTokens, "input tokens summed")
assert.Equal(t, int64(135), a.OutputTokens, "output tokens summed")
assert.Equal(t, int64(445), a.TotalTokens, "total tokens summed")
assert.InDelta(t, 0.031, a.TotalCostUSD(), 1e-9, "cost summed")
assert.InDelta(t, 0.031, a.CostUSD, 1e-9, "cost summed")
assert.Equal(t, "deny", a.Decision, "any deny makes the session a deny")
assert.ElementsMatch(t, []string{"openai", "anthropic"}, a.Providers, "distinct providers")
assert.ElementsMatch(t, []string{"gpt-5.4", "claude-haiku-4-5"}, a.Models, "distinct models")

View File

@@ -25,8 +25,8 @@ type Model struct {
// These typically need NetBird identity stamped onto upstream
// requests so the gateway's analytics and budgets attribute to the
// real caller; that's what IdentityInjection is for.
// - KindCustom: named and generic OpenAI-compatible self-hosted
// endpoints (vLLM, Ollama, custom inference servers).
// - KindCustom: the catch-all "OpenAI-compatible self-hosted endpoint"
// entry (vLLM, Ollama, custom inference servers).
//
// Frontend uses Kind to group the provider Select in the modal so an
// operator can spot at a glance which catalog entries proxy other
@@ -40,18 +40,6 @@ const (
KindCustom ProviderKind = "custom"
)
// AuthMode describes whether a catalog provider accepts an upstream
// credential. The zero value intentionally resolves to AuthModeRequired so
// existing and newly-added catalog entries fail closed unless they explicitly
// opt into a less restrictive mode.
type AuthMode string
const (
AuthModeRequired AuthMode = "required"
AuthModeOptional AuthMode = "optional"
AuthModeNone AuthMode = "none"
)
// Provider is the in-memory representation of a catalog provider.
type Provider struct {
ID string
@@ -60,9 +48,6 @@ type Provider struct {
DefaultHost string
// Kind groups this entry for UI presentation; see ProviderKind.
Kind ProviderKind
// AuthMode declares whether the provider requires, optionally accepts, or
// does not support an upstream API key. Empty defaults to required.
AuthMode AuthMode
// AuthHeaderName is the HTTP header the provider's API expects
// the credential under (e.g. "Authorization" for OpenAI,
// "x-api-key" for Anthropic). Combined with AuthHeaderTemplate
@@ -101,17 +86,6 @@ type Provider struct {
// upstream provider + credentials on Portkey's hosted side).
ExtraHeaders []ExtraHeader
Models []Model
// ModelDiscovery configures provider-specific discovery behavior. A nil
// profile means discovery is unsupported. The profile never carries
// caller-supplied paths: the proxy owns the fixed endpoint allowlist.
ModelDiscovery *ModelDiscovery
}
// ModelDiscovery describes the safe, catalog-owned fallback behavior for a
// discoverable provider. Every discovery starts with OpenAI-compatible
// /v1/models; OllamaFallback permits /api/tags only when that route is absent.
type ModelDiscovery struct {
OllamaFallback bool
}
// ExtraHeader names a single optional per-provider routing/config
@@ -727,31 +701,11 @@ var providers = []Provider{
BrandColor: "#30A2FF",
Models: []Model{},
},
{
// Ollama exposes an OpenAI-compatible /v1 API. Like vLLM, it gets a
// dedicated catalog id for provider-specific setup guidance while
// retaining generic custom-provider routing. Authentication is optional
// because local Ollama is authless but protected front ends may use it.
ID: "ollama",
Kind: KindCustom,
AuthMode: AuthModeOptional,
Name: "Ollama",
Description: "Self-hosted Ollama (OpenAI-compatible)",
DefaultHost: "",
AuthHeaderName: "Authorization",
AuthHeaderTemplate: "Bearer ${API_KEY}",
DefaultContentType: "application/json",
BrandColor: "#000000",
Models: []Model{},
ModelDiscovery: &ModelDiscovery{
OllamaFallback: true,
},
},
{
ID: "custom",
Kind: KindCustom,
Name: "Custom / Self-hosted",
Description: "Other OpenAI-compatible endpoint",
Description: "OpenAI-compatible endpoint (vLLM, Ollama, …)",
DefaultHost: "",
AuthHeaderName: "Authorization",
AuthHeaderTemplate: "Bearer ${API_KEY}",
@@ -784,15 +738,6 @@ func IsKnown(id string) bool {
return ok
}
// EffectiveAuthMode returns the provider's configured authentication mode.
// Treating the zero value as required keeps older catalog entries fail closed.
func (p Provider) EffectiveAuthMode() AuthMode {
if p.AuthMode == "" {
return AuthModeRequired
}
return p.AuthMode
}
// IsVertexPathStyle reports whether a provider uses the Google Vertex AI
// request shape — the model is carried in the URL path
// (/v1/projects/{p}/locations/{r}/publishers/{pub}/models/{model}:{action})
@@ -829,17 +774,15 @@ func (p Provider) ToAPIResponse() api.AgentNetworkCatalogProvider {
kind = api.AgentNetworkCatalogProviderKindCustom
}
resp := api.AgentNetworkCatalogProvider{
Id: p.ID,
Name: p.Name,
Description: p.Description,
DefaultHost: p.DefaultHost,
Kind: kind,
AuthMode: api.AgentNetworkCatalogProviderAuthMode(p.EffectiveAuthMode()),
AuthHeaderTemplate: p.AuthHeaderTemplate,
DefaultContentType: p.DefaultContentType,
BrandColor: p.BrandColor,
Models: models,
SupportsModelDiscovery: p.ModelDiscovery != nil,
Id: p.ID,
Name: p.Name,
Description: p.Description,
DefaultHost: p.DefaultHost,
Kind: kind,
AuthHeaderTemplate: p.AuthHeaderTemplate,
DefaultContentType: p.DefaultContentType,
BrandColor: p.BrandColor,
Models: models,
}
if len(p.ExtraHeaders) > 0 {
extras := make([]api.AgentNetworkCatalogExtraHeader, 0, len(p.ExtraHeaders))

View File

@@ -1,65 +0,0 @@
package catalog
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/shared/management/http/api"
)
func TestOllamaCatalogEntry(t *testing.T) {
entry, ok := Lookup("ollama")
require.True(t, ok, "Ollama must be available as a dedicated catalog provider")
assert.Equal(t, KindCustom, entry.Kind)
assert.Equal(t, AuthModeOptional, entry.EffectiveAuthMode())
assert.Equal(t, "Ollama", entry.Name)
assert.Equal(t, "Self-hosted Ollama (OpenAI-compatible)", entry.Description)
assert.Empty(t, entry.DefaultHost, "the dashboard owns Ollama's scheme-aware HTTP placeholder")
assert.Equal(t, "Authorization", entry.AuthHeaderName)
assert.Equal(t, "Bearer ${API_KEY}", entry.AuthHeaderTemplate)
assert.Equal(t, "application/json", entry.DefaultContentType)
assert.Empty(t, entry.ParserID, "Ollama preserves the untagged vLLM/custom routing behavior")
assert.Empty(t, entry.Models, "Ollama models are installed dynamically on the configured endpoint")
require.NotNil(t, entry.ModelDiscovery)
assert.True(t, entry.ModelDiscovery.OllamaFallback)
wire := entry.ToAPIResponse()
assert.Equal(t, "ollama", wire.Id)
assert.Equal(t, api.AgentNetworkCatalogProviderKindCustom, wire.Kind)
assert.Equal(t, api.AgentNetworkCatalogProviderAuthModeOptional, wire.AuthMode)
assert.True(t, wire.SupportsModelDiscovery)
assert.NotNil(t, wire.Models)
assert.Empty(t, wire.Models)
}
func TestOnlyOllamaSupportsModelDiscovery(t *testing.T) {
for _, entry := range All() {
supportsDiscovery := entry.ModelDiscovery != nil
assert.Equal(t, entry.ID == "ollama", supportsDiscovery, entry.ID)
assert.Equal(t, supportsDiscovery, entry.ToAPIResponse().SupportsModelDiscovery, entry.ID)
}
}
func TestCatalogAuthenticationModes(t *testing.T) {
openAI, ok := Lookup("openai_api")
require.True(t, ok)
assert.Empty(t, openAI.AuthMode, "existing entries use the fail-closed default")
assert.Equal(t, AuthModeRequired, openAI.EffectiveAuthMode())
assert.Equal(t, api.AgentNetworkCatalogProviderAuthModeRequired, openAI.ToAPIResponse().AuthMode)
for _, entry := range All() {
switch entry.EffectiveAuthMode() {
case AuthModeRequired, AuthModeOptional:
assert.NotEmpty(t, entry.AuthHeaderName, "%s must declare an auth header name", entry.ID)
assert.NotEmpty(t, entry.AuthHeaderTemplate, "%s must declare an auth header template", entry.ID)
case AuthModeNone:
assert.Empty(t, entry.AuthHeaderName, "%s must not declare an auth header name", entry.ID)
assert.Empty(t, entry.AuthHeaderTemplate, "%s must not declare an auth header template", entry.ID)
default:
t.Errorf("%s has invalid auth mode %q", entry.ID, entry.AuthMode)
}
}
}

View File

@@ -1,72 +0,0 @@
package handlers
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/gorilla/mux"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/shared/auth"
"github.com/netbirdio/netbird/shared/management/http/api"
)
type discoveryManagerStub struct {
agentnetwork.Manager
result *agentNetworkTypes.ModelDiscoveryResult
err error
accountID string
userID string
providerID string
}
func (m *discoveryManagerStub) DiscoverProviderModels(_ context.Context, accountID, userID, providerID string) (*agentNetworkTypes.ModelDiscoveryResult, error) {
m.accountID = accountID
m.userID = userID
m.providerID = providerID
return m.result, m.err
}
func TestDiscoverProviderModelsHandler(t *testing.T) {
manager := &discoveryManagerStub{
Manager: agentnetwork.NewManagerMock(),
result: &agentNetworkTypes.ModelDiscoveryResult{
RequestID: "probe-123",
Source: "ollama_api_tags",
ProxyCluster: "private.example.com",
Models: []agentNetworkTypes.DiscoveredModel{
{ID: "llama3.2:latest", Label: "llama3.2:latest"},
},
},
}
router := mux.NewRouter()
RegisterEndpoints(manager, router)
req := httptest.NewRequest(http.MethodPost, "/agent-network/providers/provider-1/discover-models", nil)
req = nbcontext.SetUserAuthInRequest(req, auth.UserAuth{
UserId: testUserID,
AccountId: testAccountID,
})
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
assert.Equal(t, "no-store", rec.Header().Get("Cache-Control"))
var response api.AgentNetworkModelDiscoveryResponse
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &response))
assert.Equal(t, testAccountID, manager.accountID)
assert.Equal(t, testUserID, manager.userID)
assert.Equal(t, "provider-1", manager.providerID)
assert.Equal(t, "probe-123", response.RequestId)
assert.Equal(t, "private.example.com", response.ProxyCluster)
assert.Equal(t, api.AgentNetworkModelDiscoveryResponseSourceOllamaApiTags, response.Source)
require.Len(t, response.Models, 1)
assert.Equal(t, "llama3.2:latest", response.Models[0].Id)
}

View File

@@ -35,7 +35,6 @@ func RegisterEndpoints(manager agentnetwork.Manager, router *mux.Router) {
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET", "OPTIONS")
router.HandleFunc("/agent-network/providers/{providerId}", h.updateProvider).Methods("PUT", "OPTIONS")
router.HandleFunc("/agent-network/providers/{providerId}", h.deleteProvider).Methods("DELETE", "OPTIONS")
router.HandleFunc("/agent-network/providers/{providerId}/discover-models", h.discoverProviderModels).Methods("POST", "OPTIONS")
h.addPolicyEndpoints(router)
h.addGuardrailEndpoints(router)
h.addSettingsEndpoints(router)
@@ -99,41 +98,6 @@ func (h *handler) getProvider(w http.ResponseWriter, r *http.Request) {
util.WriteJSONObject(r.Context(), w, provider.ToAPIResponse())
}
func (h *handler) discoverProviderModels(w http.ResponseWriter, r *http.Request) {
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
if err != nil {
util.WriteError(r.Context(), err, w)
return
}
providerID := strings.TrimSpace(mux.Vars(r)["providerId"])
if providerID == "" {
util.WriteError(r.Context(), status.Errorf(status.InvalidArgument, "provider ID is required"), w)
return
}
result, err := h.manager.DiscoverProviderModels(r.Context(), userAuth.AccountId, userAuth.UserId, providerID)
if err != nil {
util.WriteError(r.Context(), err, w)
return
}
models := make([]api.AgentNetworkDiscoveredModel, 0, len(result.Models))
for _, model := range result.Models {
models = append(models, api.AgentNetworkDiscoveredModel{
Id: model.ID,
Label: model.Label,
})
}
w.Header().Set("Cache-Control", "no-store")
util.WriteJSONObject(r.Context(), w, api.AgentNetworkModelDiscoveryResponse{
Models: models,
Source: api.AgentNetworkModelDiscoveryResponseSource(result.Source),
ProxyCluster: result.ProxyCluster,
RequestId: result.RequestID,
})
}
func (h *handler) createProvider(w http.ResponseWriter, r *http.Request) {
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
if err != nil {
@@ -229,12 +193,11 @@ func (h *handler) deleteProvider(w http.ResponseWriter, r *http.Request) {
util.WriteJSONObject(r.Context(), w, util.EmptyObject{})
}
func validate(req *api.AgentNetworkProviderRequest, creating bool) error {
func validate(req *api.AgentNetworkProviderRequest, requireAPIKey bool) error {
if strings.TrimSpace(req.ProviderId) == "" {
return status.Errorf(status.InvalidArgument, "provider_id is required")
}
entry, ok := catalog.Lookup(req.ProviderId)
if !ok {
if !catalog.IsKnown(req.ProviderId) {
return status.Errorf(status.InvalidArgument, "provider_id %q is not a known catalog provider", req.ProviderId)
}
if strings.TrimSpace(req.Name) == "" {
@@ -247,27 +210,8 @@ func validate(req *api.AgentNetworkProviderRequest, creating bool) error {
if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") {
return status.Errorf(status.InvalidArgument, "upstream_url must be a full http(s) URL")
}
return validateProviderAPIKey(entry, req.ApiKey, creating)
}
func validateProviderAPIKey(entry catalog.Provider, apiKey *string, creating bool) error {
apiKeyBlank := apiKey == nil || strings.TrimSpace(*apiKey) == ""
switch entry.EffectiveAuthMode() {
case catalog.AuthModeRequired:
// Updates may omit the key to preserve it, but an explicitly empty
// value cannot clear a credential required by the selected provider.
if (creating || apiKey != nil) && apiKeyBlank {
return status.Errorf(status.InvalidArgument, "api_key is required for provider_id %q", entry.ID)
}
case catalog.AuthModeOptional:
// Both omitted and explicitly empty values are valid. The manager
// distinguishes preserve from clear during update.
case catalog.AuthModeNone:
if !apiKeyBlank {
return status.Errorf(status.InvalidArgument, "api_key is not supported for provider_id %q", entry.ID)
}
default:
return status.Errorf(status.InvalidArgument, "provider_id %q has an invalid authentication mode", entry.ID)
if requireAPIKey && (req.ApiKey == nil || strings.TrimSpace(*req.ApiKey) == "") {
return status.Errorf(status.InvalidArgument, "api_key is required")
}
return nil
}

View File

@@ -1,47 +0,0 @@
package handlers
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
)
func TestValidateProviderAPIKey(t *testing.T) {
key := "secret"
empty := ""
tests := []struct {
name string
mode catalog.AuthMode
apiKey *string
creating bool
wantErr string
}{
{name: "required create with key", mode: catalog.AuthModeRequired, apiKey: &key, creating: true},
{name: "required create omitted", mode: catalog.AuthModeRequired, creating: true, wantErr: "required"},
{name: "required update omitted preserves", mode: catalog.AuthModeRequired},
{name: "required update explicit empty", mode: catalog.AuthModeRequired, apiKey: &empty, wantErr: "required"},
{name: "optional create omitted", mode: catalog.AuthModeOptional, creating: true},
{name: "optional update explicit empty", mode: catalog.AuthModeOptional, apiKey: &empty},
{name: "optional with key", mode: catalog.AuthModeOptional, apiKey: &key, creating: true},
{name: "none omitted", mode: catalog.AuthModeNone, creating: true},
{name: "none explicit empty", mode: catalog.AuthModeNone, apiKey: &empty},
{name: "none with key", mode: catalog.AuthModeNone, apiKey: &key, wantErr: "not supported"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
entry := catalog.Provider{ID: "test-provider", AuthMode: tt.mode}
err := validateProviderAPIKey(entry, tt.apiKey, tt.creating)
if tt.wantErr == "" {
require.NoError(t, err)
return
}
require.Error(t, err)
assert.Contains(t, err.Error(), tt.wantErr)
})
}
}

View File

@@ -9,12 +9,9 @@ import (
"strings"
"sync"
"time"
"unicode"
"unicode/utf8"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/catalog"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/labelgen"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
@@ -51,7 +48,6 @@ func ensureSessionKeys(p *types.Provider) error {
type Manager interface {
GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error)
GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error)
DiscoverProviderModels(ctx context.Context, accountID, userID, providerID string) (*types.ModelDiscoveryResult, error)
CreateProvider(ctx context.Context, userID string, provider *types.Provider, bootstrapCluster string) (*types.Provider, error)
UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
DeleteProvider(ctx context.Context, accountID, userID, providerID string) error
@@ -132,27 +128,8 @@ type managerImpl struct {
// state; concurrent provider creates would otherwise race.
labelRngMu sync.Mutex
labelRng *rand.Rand
// discoveryAttempts provides lightweight per-provider admission control for
// the explicit endpoint probe. It prevents duplicate clicks from occupying
// multiple proxy control-stream requests at once and adds a short cooldown
// after each attempt.
discoveryMu sync.Mutex
discoveryAttempts map[string]modelDiscoveryAttempt
}
type modelDiscoveryAttempt struct {
inFlight bool
lastStarted time.Time
}
const (
modelDiscoveryTimeout = 10 * time.Second
modelDiscoveryCooldown = 2 * time.Second
maxDiscoveredModels = 500
maxDiscoveredModelLen = 512
)
// NewManager constructs the persistent Agent Network manager. The
// manager persists provider/policy/guardrail configuration and, on
// every mutation, reconciles the in-memory synthesised reverse-proxy
@@ -171,7 +148,6 @@ func NewManager(
proxyController: proxyController,
reconcileCache: make(map[string]map[string]*proto.ProxyMapping),
labelRng: rand.New(rand.NewSource(time.Now().UnixNano())),
discoveryAttempts: make(map[string]modelDiscoveryAttempt),
}
}
@@ -189,202 +165,6 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid
return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
}
// DiscoverProviderModels asks a capable proxy in the account's selected
// cluster to query the persisted provider endpoint. The browser supplies only
// the provider id: URL, TLS policy, and credential are loaded here so a caller
// cannot turn this operation into an arbitrary network probe.
func (m *managerImpl) DiscoverProviderModels(ctx context.Context, accountID, userID, providerID string) (*types.ModelDiscoveryResult, error) {
if err := m.requirePermission(ctx, accountID, userID, operations.Update); err != nil {
return nil, err
}
providerID = strings.TrimSpace(providerID)
if providerID == "" {
return nil, status.Errorf(status.InvalidArgument, "provider ID is required")
}
provider, err := m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
if err != nil {
return nil, err
}
entry, ok := catalog.Lookup(provider.ProviderID)
if !ok {
return nil, status.Errorf(status.PreconditionFailed, "provider references an unknown catalog provider")
}
if entry.ModelDiscovery == nil {
return nil, status.Errorf(status.PreconditionFailed, "provider type %q does not support model discovery", provider.ProviderID)
}
if m.proxyController == nil {
return nil, status.Errorf(status.PreconditionFailed, "model discovery is unavailable")
}
settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
if err != nil {
var statusErr *status.Error
if errors.As(err, &statusErr) && statusErr.Type() == status.NotFound {
return nil, status.Errorf(status.PreconditionFailed, "configure an Agent Network proxy cluster before discovering models")
}
return nil, err
}
cluster := strings.TrimSpace(settings.Cluster)
if cluster == "" {
return nil, status.Errorf(status.PreconditionFailed, "configure an Agent Network proxy cluster before discovering models")
}
authHeaderName, authHeaderValue, gcpKey, err := providerAuthHeader(provider)
if err != nil {
return nil, status.Errorf(status.PreconditionFailed, "provider authentication is not configured correctly")
}
if gcpKey != "" {
return nil, status.Errorf(status.PreconditionFailed, "provider authentication is not supported for model discovery")
}
attemptKey := accountID + "\x00" + provider.ID
if !m.beginModelDiscovery(attemptKey, time.Now()) {
return nil, status.Errorf(status.TooManyRequests, "model discovery is already running or was requested too recently")
}
defer m.finishModelDiscovery(attemptKey)
probeCtx, cancel := context.WithTimeout(ctx, modelDiscoveryTimeout)
defer cancel()
probeResult, err := m.proxyController.DiscoverModels(probeCtx, accountID, cluster, &proto.ModelDiscoveryRequest{
UpstreamUrl: strings.TrimSpace(provider.UpstreamURL),
AuthHeaderName: authHeaderName,
AuthHeaderValue: authHeaderValue,
SkipTlsVerify: provider.SkipTLSVerification,
OllamaFallback: entry.ModelDiscovery.OllamaFallback,
})
if err != nil {
if errors.Is(err, context.DeadlineExceeded) || errors.Is(probeCtx.Err(), context.DeadlineExceeded) {
return nil, status.Errorf(status.PreconditionFailed, "model discovery timed out")
}
if _, ok := status.FromError(err); ok {
return nil, err
}
return nil, status.Errorf(status.PreconditionFailed, "model discovery could not be completed")
}
if probeResult == nil {
return nil, status.Errorf(status.PreconditionFailed, "proxy returned an empty model discovery response")
}
if strings.TrimSpace(probeResult.Error) != "" {
message := safeModelDiscoveryText(probeResult.Error, 256)
if message == "" {
message = "proxy reported a discovery failure"
}
requestID := safeModelDiscoveryText(probeResult.RequestId, 128)
if requestID != "" {
return nil, status.Errorf(status.PreconditionFailed, "model discovery failed: %s (request_id: %s)", message, requestID)
}
return nil, status.Errorf(status.PreconditionFailed, "model discovery failed: %s", message)
}
requestID := safeModelDiscoveryText(probeResult.RequestId, 128)
if requestID == "" {
return nil, status.Errorf(status.PreconditionFailed, "proxy returned an uncorrelated model discovery response")
}
source := strings.TrimSpace(probeResult.Source)
switch source {
case "openai_v1_models", "ollama_api_tags":
default:
return nil, status.Errorf(status.PreconditionFailed, "proxy returned an unsupported model discovery response")
}
return &types.ModelDiscoveryResult{
RequestID: requestID,
Source: source,
ProxyCluster: cluster,
Models: normalizeDiscoveredModels(probeResult.Models),
}, nil
}
func (m *managerImpl) beginModelDiscovery(key string, now time.Time) bool {
m.discoveryMu.Lock()
defer m.discoveryMu.Unlock()
if m.discoveryAttempts == nil {
m.discoveryAttempts = make(map[string]modelDiscoveryAttempt)
}
attempt := m.discoveryAttempts[key]
if attempt.inFlight || (!attempt.lastStarted.IsZero() && now.Sub(attempt.lastStarted) < modelDiscoveryCooldown) {
return false
}
attempt.inFlight = true
attempt.lastStarted = now
m.discoveryAttempts[key] = attempt
return true
}
func (m *managerImpl) finishModelDiscovery(key string) {
m.discoveryMu.Lock()
attempt, ok := m.discoveryAttempts[key]
if !ok {
m.discoveryMu.Unlock()
return
}
attempt.inFlight = false
m.discoveryAttempts[key] = attempt
m.discoveryMu.Unlock()
remaining := time.Until(attempt.lastStarted.Add(modelDiscoveryCooldown))
if remaining <= 0 {
m.expireModelDiscoveryAttempt(key, attempt.lastStarted)
return
}
time.AfterFunc(remaining, func() {
m.expireModelDiscoveryAttempt(key, attempt.lastStarted)
})
}
func (m *managerImpl) expireModelDiscoveryAttempt(key string, lastStarted time.Time) {
m.discoveryMu.Lock()
defer m.discoveryMu.Unlock()
attempt, ok := m.discoveryAttempts[key]
if ok && !attempt.inFlight && attempt.lastStarted.Equal(lastStarted) {
delete(m.discoveryAttempts, key)
}
}
func normalizeDiscoveredModels(models []*proto.ModelDiscoveryModel) []types.DiscoveredModel {
out := make([]types.DiscoveredModel, 0, min(len(models), maxDiscoveredModels))
seen := make(map[string]struct{}, min(len(models), maxDiscoveredModels))
for _, model := range models {
if model == nil {
continue
}
id := safeModelDiscoveryText(model.Id, maxDiscoveredModelLen)
if id == "" || len(id) > maxDiscoveredModelLen {
continue
}
if _, ok := seen[id]; ok {
continue
}
label := safeModelDiscoveryText(model.Label, maxDiscoveredModelLen)
if label == "" || len(label) > maxDiscoveredModelLen {
label = id
}
seen[id] = struct{}{}
out = append(out, types.DiscoveredModel{ID: id, Label: label})
if len(out) == maxDiscoveredModels {
break
}
}
slices.SortFunc(out, func(a, b types.DiscoveredModel) int {
return strings.Compare(a.ID, b.ID)
})
return out
}
func safeModelDiscoveryText(value string, maxLen int) string {
value = strings.TrimSpace(value)
if value == "" || len(value) > maxLen || !utf8.ValidString(value) {
return ""
}
for _, r := range value {
if unicode.IsControl(r) {
return ""
}
}
return value
}
// CreateProvider persists a new provider for the account. bootstrapCluster
// is used only when the per-account agent-network Settings row hasn't
// been created yet; otherwise it is ignored (the cluster is pinned on
@@ -394,8 +174,11 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide
return nil, err
}
if err := prepareProviderAPIKey(provider, nil); err != nil {
return nil, err
// An empty api_key would silently produce a synthesised service
// that 401s on every upstream request. Surface the misconfiguration
// at create time instead.
if strings.TrimSpace(provider.APIKey) == "" {
return nil, status.Errorf(status.InvalidArgument, "api_key is required when creating an agent network provider")
}
if provider.ID == "" {
@@ -439,8 +222,13 @@ func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provide
return nil, fmt.Errorf("failed to get agent network provider: %w", err)
}
if err := prepareProviderAPIKey(provider, existing); err != nil {
return nil, err
// Preserve the API key if the caller didn't rotate it. A
// whitespace-only value is treated as "not rotated" rather than a
// real key, but it must not silently overwrite a valid stored key.
if provider.APIKey == "" {
provider.APIKey = existing.APIKey
} else if strings.TrimSpace(provider.APIKey) == "" {
return nil, status.Errorf(status.InvalidArgument, "api_key must be non-blank when rotating an agent network provider")
}
// Always preserve the session keypair across updates so existing
// session cookies stay valid. The keys are server-managed and
@@ -463,52 +251,6 @@ func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provide
return provider, nil
}
// prepareProviderAPIKey applies catalog-owned authentication semantics before
// a provider is persisted. existing is nil on create. On update, omission
// preserves a key only when the provider type is unchanged; switching types
// never carries an old provider's secret into the new upstream.
func prepareProviderAPIKey(provider, existing *types.Provider) error {
entry, ok := catalog.Lookup(provider.ProviderID)
if !ok {
return status.Errorf(status.InvalidArgument, "provider_id %q is not a known catalog provider", provider.ProviderID)
}
return prepareProviderAPIKeyForEntry(provider, existing, entry)
}
func prepareProviderAPIKeyForEntry(provider, existing *types.Provider, entry catalog.Provider) error {
keyProvided := provider.APIKeyProvided || provider.APIKey != ""
providerChanged := existing != nil && provider.ProviderID != existing.ProviderID
if existing != nil && !keyProvided && !providerChanged && entry.EffectiveAuthMode() != catalog.AuthModeNone {
provider.APIKey = existing.APIKey
}
// Normalize after a preserved value is restored so a legacy
// whitespace-only credential cannot pass manager validation and then fail
// later during synthesis.
if strings.TrimSpace(provider.APIKey) == "" {
provider.APIKey = ""
}
switch entry.EffectiveAuthMode() {
case catalog.AuthModeRequired:
if provider.APIKey == "" {
return status.Errorf(status.InvalidArgument, "api_key is required for provider_id %q", provider.ProviderID)
}
case catalog.AuthModeOptional:
// Empty is a valid credential state.
case catalog.AuthModeNone:
if keyProvided && provider.APIKey != "" {
return status.Errorf(status.InvalidArgument, "api_key is not supported for provider_id %q", provider.ProviderID)
}
// Clear any stale credential already stored for a provider whose
// catalog semantics changed to no authentication.
provider.APIKey = ""
default:
return status.Errorf(status.InvalidArgument, "provider_id %q has an invalid authentication mode", provider.ProviderID)
}
return nil
}
func (m *managerImpl) DeleteProvider(ctx context.Context, accountID, userID, providerID string) error {
if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil {
return err
@@ -1066,10 +808,6 @@ func (*mockManager) GetProvider(_ context.Context, _, _, _ string) (*types.Provi
return &types.Provider{}, nil
}
func (*mockManager) DiscoverProviderModels(_ context.Context, _, _, _ string) (*types.ModelDiscoveryResult, error) {
return nil, status.Errorf(status.PreconditionFailed, "model discovery is unavailable")
}
func (*mockManager) CreateProvider(_ context.Context, _ string, p *types.Provider, _ string) (*types.Provider, error) {
return p, nil
}

View File

@@ -1,180 +0,0 @@
package agentnetwork
import (
"context"
"strings"
"testing"
"time"
"github.com/golang/mock/gomock"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/management/status"
)
func newModelDiscoveryManager(t *testing.T) (*managerImpl, *store.MockStore, *permissions.MockManager, *proxy.MockController) {
t.Helper()
ctrl := gomock.NewController(t)
mockStore := store.NewMockStore(ctrl)
mockPermissions := permissions.NewMockManager(ctrl)
mockProxy := proxy.NewMockController(ctrl)
return &managerImpl{
store: mockStore,
permissionsManager: mockPermissions,
proxyController: mockProxy,
discoveryAttempts: make(map[string]modelDiscoveryAttempt),
}, mockStore, mockPermissions, mockProxy
}
func allowModelDiscovery(mockPermissions *permissions.MockManager) {
mockPermissions.EXPECT().
ValidateUserPermissions(gomock.Any(), "account-1", "user-1", modules.AgentNetwork, operations.Update).
Return(true, context.Background(), nil)
}
func TestDiscoverProviderModelsUsesPersistedProviderAndCluster(t *testing.T) {
manager, mockStore, mockPermissions, mockProxy := newModelDiscoveryManager(t)
allowModelDiscovery(mockPermissions)
provider := &types.Provider{
ID: "provider-1",
AccountID: "account-1",
ProviderID: "ollama",
UpstreamURL: "http://ollama.internal:11434/base",
APIKey: "stored-key",
SkipTLSVerification: true,
}
mockStore.EXPECT().
GetAgentNetworkProviderByID(gomock.Any(), store.LockingStrengthNone, "account-1", "provider-1").
Return(provider, nil)
mockStore.EXPECT().
GetAgentNetworkSettings(gomock.Any(), store.LockingStrengthNone, "account-1").
Return(&types.Settings{AccountID: "account-1", Cluster: "private.example.com"}, nil)
mockProxy.EXPECT().
DiscoverModels(gomock.Any(), "account-1", "private.example.com", gomock.Any()).
DoAndReturn(func(_ context.Context, _, _ string, request *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
assert.Empty(t, request.RequestId, "the control-plane server owns request IDs")
assert.Equal(t, provider.UpstreamURL, request.UpstreamUrl)
assert.Equal(t, "Authorization", request.AuthHeaderName)
assert.Equal(t, "Bearer stored-key", request.AuthHeaderValue)
assert.True(t, request.SkipTlsVerify)
assert.True(t, request.OllamaFallback)
return &proto.ModelDiscoveryResult{
RequestId: "probe-123",
Source: "openai_v1_models",
Models: []*proto.ModelDiscoveryModel{
{Id: "zeta", Label: ""},
{Id: "alpha", Label: "Alpha"},
{Id: "zeta", Label: "duplicate"},
{Id: "bad\x00model", Label: "bad"},
{Id: strings.Repeat("x", maxDiscoveredModelLen+1), Label: "too long"},
},
}, nil
})
result, err := manager.DiscoverProviderModels(context.Background(), "account-1", "user-1", "provider-1")
require.NoError(t, err)
assert.Equal(t, "probe-123", result.RequestID)
assert.Equal(t, "openai_v1_models", result.Source)
assert.Equal(t, "private.example.com", result.ProxyCluster)
assert.Equal(t, []types.DiscoveredModel{
{ID: "alpha", Label: "Alpha"},
{ID: "zeta", Label: "zeta"},
}, result.Models)
}
func TestDiscoverProviderModelsRejectsUnsupportedProvider(t *testing.T) {
manager, mockStore, mockPermissions, _ := newModelDiscoveryManager(t)
allowModelDiscovery(mockPermissions)
mockStore.EXPECT().
GetAgentNetworkProviderByID(gomock.Any(), store.LockingStrengthNone, "account-1", "provider-1").
Return(&types.Provider{
ID: "provider-1",
AccountID: "account-1",
ProviderID: "openai_api",
}, nil)
_, err := manager.DiscoverProviderModels(context.Background(), "account-1", "user-1", "provider-1")
require.Error(t, err)
statusErr, ok := status.FromError(err)
require.True(t, ok)
assert.Equal(t, status.PreconditionFailed, statusErr.Type())
assert.Contains(t, err.Error(), "does not support")
}
func TestDiscoverProviderModelsPreservesCorrelatedProxyError(t *testing.T) {
manager, mockStore, mockPermissions, mockProxy := newModelDiscoveryManager(t)
allowModelDiscovery(mockPermissions)
mockStore.EXPECT().
GetAgentNetworkProviderByID(gomock.Any(), store.LockingStrengthNone, "account-1", "provider-1").
Return(&types.Provider{
ID: "provider-1",
AccountID: "account-1",
ProviderID: "ollama",
UpstreamURL: "http://ollama.internal:11434",
}, nil)
mockStore.EXPECT().
GetAgentNetworkSettings(gomock.Any(), store.LockingStrengthNone, "account-1").
Return(&types.Settings{AccountID: "account-1", Cluster: "private.example.com"}, nil)
mockProxy.EXPECT().
DiscoverModels(gomock.Any(), "account-1", "private.example.com", gomock.Any()).
Return(&proto.ModelDiscoveryResult{
RequestId: "probe-failed",
Error: "upstream returned HTTP 401",
}, nil)
_, err := manager.DiscoverProviderModels(context.Background(), "account-1", "user-1", "provider-1")
require.Error(t, err)
assert.Contains(t, err.Error(), "upstream returned HTTP 401")
assert.Contains(t, err.Error(), "probe-failed")
}
func TestModelDiscoveryAdmissionControl(t *testing.T) {
manager := &managerImpl{}
start := time.Now()
require.True(t, manager.beginModelDiscovery("account/provider", start))
assert.False(t, manager.beginModelDiscovery("account/provider", start.Add(time.Second)))
manager.finishModelDiscovery("account/provider")
assert.False(t, manager.beginModelDiscovery("account/provider", start.Add(time.Second)))
assert.True(t, manager.beginModelDiscovery("account/provider", start.Add(modelDiscoveryCooldown)))
}
func TestModelDiscoveryAdmissionControlExpiresFinishedAttempts(t *testing.T) {
manager := &managerImpl{}
start := time.Now()
require.True(t, manager.beginModelDiscovery("account/provider", start))
manager.finishModelDiscovery("account/provider")
manager.expireModelDiscoveryAttempt("account/provider", start)
manager.discoveryMu.Lock()
_, retained := manager.discoveryAttempts["account/provider"]
manager.discoveryMu.Unlock()
assert.False(t, retained)
}
func TestModelDiscoveryAdmissionControlKeepsReplacementAttempt(t *testing.T) {
manager := &managerImpl{}
start := time.Now()
require.True(t, manager.beginModelDiscovery("account/provider", start))
manager.finishModelDiscovery("account/provider")
require.True(t, manager.beginModelDiscovery("account/provider", start.Add(modelDiscoveryCooldown)))
manager.expireModelDiscoveryAttempt("account/provider", start)
manager.discoveryMu.Lock()
attempt, retained := manager.discoveryAttempts["account/provider"]
manager.discoveryMu.Unlock()
require.True(t, retained)
assert.True(t, attempt.inFlight)
assert.Equal(t, start.Add(modelDiscoveryCooldown), attempt.lastStarted)
}

View File

@@ -1,107 +0,0 @@
package agentnetwork
import (
"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/types"
)
func TestPrepareProviderAPIKey(t *testing.T) {
t.Run("required create rejects an empty key", func(t *testing.T) {
provider := &types.Provider{ProviderID: "openai_api"}
err := prepareProviderAPIKey(provider, nil)
require.Error(t, err)
assert.Contains(t, err.Error(), "api_key is required")
})
t.Run("optional create accepts an empty key", func(t *testing.T) {
provider := &types.Provider{ProviderID: "ollama"}
require.NoError(t, prepareProviderAPIKey(provider, nil))
assert.Empty(t, provider.APIKey)
})
t.Run("optional update omission preserves the existing key", func(t *testing.T) {
existing := &types.Provider{ProviderID: "ollama", APIKey: "existing"}
provider := &types.Provider{ProviderID: "ollama"}
require.NoError(t, prepareProviderAPIKey(provider, existing))
assert.Equal(t, "existing", provider.APIKey)
})
t.Run("optional update explicit empty clears the existing key", func(t *testing.T) {
existing := &types.Provider{ProviderID: "ollama", APIKey: "existing"}
provider := &types.Provider{ProviderID: "ollama", APIKeyProvided: true}
require.NoError(t, prepareProviderAPIKey(provider, existing))
assert.Empty(t, provider.APIKey)
})
t.Run("changing to optional without a key does not carry the old secret", func(t *testing.T) {
existing := &types.Provider{ProviderID: "openai_api", APIKey: "sk-old-provider"}
provider := &types.Provider{ProviderID: "ollama"}
require.NoError(t, prepareProviderAPIKey(provider, existing))
assert.Empty(t, provider.APIKey)
})
t.Run("changing to required without a key fails", func(t *testing.T) {
existing := &types.Provider{ProviderID: "ollama", APIKey: "old-optional-key"}
provider := &types.Provider{ProviderID: "openai_api"}
err := prepareProviderAPIKey(provider, existing)
require.Error(t, err)
assert.Contains(t, err.Error(), "api_key is required")
assert.Empty(t, provider.APIKey)
})
t.Run("required update omission preserves the existing key", func(t *testing.T) {
existing := &types.Provider{ProviderID: "openai_api", APIKey: "sk-existing"}
provider := &types.Provider{ProviderID: "openai_api"}
require.NoError(t, prepareProviderAPIKey(provider, existing))
assert.Equal(t, "sk-existing", provider.APIKey)
})
t.Run("required update rejects a preserved whitespace-only key", func(t *testing.T) {
existing := &types.Provider{ProviderID: "openai_api", APIKey: " "}
provider := &types.Provider{ProviderID: "openai_api"}
err := prepareProviderAPIKey(provider, existing)
require.Error(t, err)
assert.Contains(t, err.Error(), "api_key is required")
assert.Empty(t, provider.APIKey)
})
t.Run("required update explicit empty is rejected", func(t *testing.T) {
existing := &types.Provider{ProviderID: "openai_api", APIKey: "sk-existing"}
provider := &types.Provider{ProviderID: "openai_api", APIKeyProvided: true}
err := prepareProviderAPIKey(provider, existing)
require.Error(t, err)
assert.Contains(t, err.Error(), "api_key is required")
})
t.Run("unknown catalog provider is rejected", func(t *testing.T) {
provider := &types.Provider{ProviderID: "unknown", APIKey: "secret"}
err := prepareProviderAPIKey(provider, nil)
require.Error(t, err)
assert.Contains(t, err.Error(), "not a known catalog provider")
})
t.Run("none mode clears a stale key", func(t *testing.T) {
entry := catalog.Provider{ID: "authless", AuthMode: catalog.AuthModeNone}
existing := &types.Provider{ProviderID: "authless", APIKey: "stale-secret"}
provider := &types.Provider{ProviderID: "authless"}
require.NoError(t, prepareProviderAPIKeyForEntry(provider, existing, entry))
assert.Empty(t, provider.APIKey)
})
t.Run("none mode rejects a supplied key", func(t *testing.T) {
entry := catalog.Provider{ID: "authless", AuthMode: catalog.AuthModeNone}
provider := &types.Provider{
ProviderID: "authless",
APIKey: "unexpected-secret",
APIKeyProvided: true,
}
err := prepareProviderAPIKeyForEntry(provider, nil, entry)
require.Error(t, err)
assert.Contains(t, err.Error(), "not supported")
})
}

View File

@@ -906,33 +906,23 @@ const (
noopUpstreamPort = uint16(443)
)
// providerAuthHeader builds the upstream auth header pair for a provider from
// its catalog entry. Optional providers with no key and providers using no
// authentication return an empty pair. Otherwise, the synthesiser substitutes
// the decrypted API key into the catalog template and returns the pair the
// router injects after stripping inbound vendor credentials.
// providerAuthHeader builds the upstream auth header pair for a
// provider from its catalog entry. The catalog declares which header
// name and template a provider's API expects; the synthesiser
// substitutes the provider's decrypted API key into the template and
// returns the (name, value) pair the router middleware injects after
// stripping the inbound vendor auth headers.
func providerAuthHeader(p *types.Provider) (name, value, gcpSAKeyB64 string, err error) {
entry, ok := catalog.Lookup(p.ProviderID)
if !ok {
return "", "", "", fmt.Errorf("provider %s references unknown catalog id %q", p.ID, p.ProviderID)
}
switch entry.EffectiveAuthMode() {
case catalog.AuthModeNone:
return "", "", "", nil
case catalog.AuthModeOptional:
if strings.TrimSpace(p.APIKey) == "" {
return "", "", "", nil
}
case catalog.AuthModeRequired:
if strings.TrimSpace(p.APIKey) == "" {
return "", "", "", fmt.Errorf("provider %s has no api key", p.ID)
}
default:
return "", "", "", fmt.Errorf("catalog entry %q has invalid authentication mode %q", p.ProviderID, entry.AuthMode)
}
if entry.AuthHeaderName == "" || entry.AuthHeaderTemplate == "" {
return "", "", "", fmt.Errorf("catalog entry %q has no auth header configured", p.ProviderID)
}
if p.APIKey == "" {
return "", "", "", fmt.Errorf("provider %s has no api key", p.ID)
}
// A "keyfile::<base64 json>" api_key is a GCP service-account key, not a
// static bearer. The proxy mints + refreshes a short-lived OAuth token from
// it at request time, so carry the key material on the route and emit no

View File

@@ -1219,59 +1219,3 @@ 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")
}
func TestSynthesizeServices_OllamaOptionalAPIKey(t *testing.T) {
tests := []struct {
name string
apiKey string
wantHeaderName string
wantHeaderValue string
}{
{
name: "no key emits no replacement auth header",
},
{
name: "configured key emits bearer auth",
apiKey: "protected-endpoint-token",
wantHeaderName: "Authorization",
wantHeaderValue: "Bearer protected-endpoint-token",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ctx := context.Background()
ctrl := gomock.NewController(t)
defer ctrl.Finish()
mockStore := store.NewMockStore(ctrl)
provider := newSynthTestProvider()
provider.ProviderID = "ollama"
provider.Name = "Ollama"
provider.UpstreamURL = "http://ollama.internal:11434"
provider.APIKey = tt.apiKey
provider.Models = []types.ProviderModel{}
policy := newSynthTestPolicy(provider.ID, "grp-eng", "")
expectSynthBaseInputs(mockStore, ctx, newSynthTestSettings(),
[]*types.Provider{provider},
[]*types.Policy{policy},
[]*types.Guardrail{})
services, err := SynthesizeServices(ctx, mockStore, testAccountID)
require.NoError(t, err)
require.Len(t, services, 1)
var routerCfg routerConfig
for _, middleware := range services[0].Targets[0].Options.Middlewares {
if middleware.ID == middlewareIDLLMRouter {
require.NoError(t, json.Unmarshal(middleware.ConfigJSON, &routerCfg))
break
}
}
require.Len(t, routerCfg.Providers, 1)
assert.Equal(t, tt.wantHeaderName, routerCfg.Providers[0].AuthHeaderName)
assert.Equal(t, tt.wantHeaderValue, routerCfg.Providers[0].AuthHeaderValue)
})
}
}

View File

@@ -41,24 +41,8 @@ type AgentNetworkAccessLog struct {
InputTokens int64
OutputTokens int64
TotalTokens int64
// Prompt-cache buckets: read + write token counts.
CachedInputTokens int64
CacheCreationTokens int64
// Per-bucket cost breakdown — one column per token bucket the provider
// bills separately. These four are the only cost state stored: the total
// and the cache portion are derived on read (TotalCostUSD / CacheCostUSD)
// rather than stored alongside, so a stored aggregate can never drift out
// of step with the components it summarises.
//
// default:0 matters on upgrade: these columns are ALTER TABLE ADD COLUMN
// on an existing table, and without it every historical row holds NULL —
// which a raw SUM()/scan into float64 can't read. The default backfills
// them as 0, so pre-upgrade rows report an unknown split, not an error.
InputCostUSD float64 `gorm:"not null;default:0"`
CachedInputCostUSD float64 `gorm:"not null;default:0"`
CacheCreationCostUSD float64 `gorm:"not null;default:0"`
OutputCostUSD float64 `gorm:"not null;default:0"`
Stream bool
CostUSD float64
Stream bool
// Prompt capture. Only populated when prompt collection is enabled
// (account master switch AND policy guardrail). Heavy free text.
@@ -76,44 +60,19 @@ type AgentNetworkAccessLog struct {
// the reverse-proxy AccessLogEntry table.
func (AgentNetworkAccessLog) TableName() string { return "agent_network_access_log" }
// CostUSDSQLExpr is the SQL sum of the per-bucket cost columns — the total cost
// of a row. Used wherever a query has to sort or aggregate on total cost now
// that no cost_usd column is stored. Plain arithmetic over NOT NULL columns, so
// it stays portable across SQLite and Postgres.
const CostUSDSQLExpr = "(input_cost_usd + cached_input_cost_usd + cache_creation_cost_usd + output_cost_usd)"
// TotalCostUSD is the request's total cost: the sum of the four per-bucket
// costs. Derived rather than stored so it cannot disagree with the breakdown.
func (a *AgentNetworkAccessLog) TotalCostUSD() float64 {
return a.InputCostUSD + a.CachedInputCostUSD + a.CacheCreationCostUSD + a.OutputCostUSD
}
// CacheCostUSD is the portion of the total billed for prompt-cache buckets:
// cache reads plus cache writes.
func (a *AgentNetworkAccessLog) CacheCostUSD() float64 {
return a.CachedInputCostUSD + a.CacheCreationCostUSD
}
// ToAPIResponse renders the flattened entry as the API representation.
func (a *AgentNetworkAccessLog) ToAPIResponse() api.AgentNetworkAccessLog {
out := api.AgentNetworkAccessLog{
Id: a.ID,
ServiceId: a.ServiceID,
Timestamp: a.Timestamp,
StatusCode: a.StatusCode,
DurationMs: int(a.Duration.Milliseconds()),
InputTokens: a.InputTokens,
OutputTokens: a.OutputTokens,
TotalTokens: a.TotalTokens,
CachedInputTokens: a.CachedInputTokens,
CacheCreationTokens: a.CacheCreationTokens,
InputCostUsd: a.InputCostUSD,
CachedInputCostUsd: a.CachedInputCostUSD,
CacheCreationCostUsd: a.CacheCreationCostUSD,
OutputCostUsd: a.OutputCostUSD,
CostUsd: a.TotalCostUSD(),
CacheCostUsd: a.CacheCostUSD(),
Stream: &a.Stream,
Id: a.ID,
ServiceId: a.ServiceID,
Timestamp: a.Timestamp,
StatusCode: a.StatusCode,
DurationMs: int(a.Duration.Milliseconds()),
InputTokens: a.InputTokens,
OutputTokens: a.OutputTokens,
TotalTokens: a.TotalTokens,
CostUsd: a.CostUSD,
Stream: &a.Stream,
}
out.UserId = strPtr(a.UserID)
@@ -153,36 +112,20 @@ func strPtr(s string) *string {
// summary plus its ordered entries. Assembled in Go from a page of entries — it
// is not a stored table.
type AgentNetworkAccessLogSession struct {
SessionID string // empty for a session-less (singleton) request
UserID string
GroupIDs []string // union of the entries' authorising groups
StartedAt time.Time
EndedAt time.Time
RequestCount int
InputTokens int64
OutputTokens int64
TotalTokens int64
CachedInputTokens int64
CacheCreationTokens int64
InputCostUSD float64
CachedInputCostUSD float64
CacheCreationCostUSD float64
OutputCostUSD float64
Providers []string // distinct vendors seen in the session
Models []string // distinct models seen in the session
Decision string // "deny" if any entry was denied, else "allow"
Entries []*AgentNetworkAccessLog
}
// TotalCostUSD is the session's total cost: the sum of the four per-bucket
// costs accumulated across its entries.
func (sess *AgentNetworkAccessLogSession) TotalCostUSD() float64 {
return sess.InputCostUSD + sess.CachedInputCostUSD + sess.CacheCreationCostUSD + sess.OutputCostUSD
}
// CacheCostUSD is the session's prompt-cache spend: cache reads plus writes.
func (sess *AgentNetworkAccessLogSession) CacheCostUSD() float64 {
return sess.CachedInputCostUSD + sess.CacheCreationCostUSD
SessionID string // empty for a session-less (singleton) request
UserID string
GroupIDs []string // union of the entries' authorising groups
StartedAt time.Time
EndedAt time.Time
RequestCount int
InputTokens int64
OutputTokens int64
TotalTokens int64
CostUSD float64
Providers []string // distinct vendors seen in the session
Models []string // distinct models seen in the session
Decision string // "deny" if any entry was denied, else "allow"
Entries []*AgentNetworkAccessLog
}
// sessionKey is the grouping key for an entry: its session id, or — when the
@@ -262,12 +205,7 @@ func (sess *AgentNetworkAccessLogSession) foldEntry(sk *sessionSeen, e *AgentNet
sess.InputTokens += e.InputTokens
sess.OutputTokens += e.OutputTokens
sess.TotalTokens += e.TotalTokens
sess.CachedInputTokens += e.CachedInputTokens
sess.CacheCreationTokens += e.CacheCreationTokens
sess.InputCostUSD += e.InputCostUSD
sess.CachedInputCostUSD += e.CachedInputCostUSD
sess.CacheCreationCostUSD += e.CacheCreationCostUSD
sess.OutputCostUSD += e.OutputCostUSD
sess.CostUSD += e.CostUSD
if e.Timestamp.Before(sess.StartedAt) {
sess.StartedAt = e.Timestamp
}
@@ -310,22 +248,15 @@ func (sess *AgentNetworkAccessLogSession) ToAPIResponse() api.AgentNetworkAccess
}
out := api.AgentNetworkAccessLogSession{
StartedAt: sess.StartedAt,
EndedAt: sess.EndedAt,
RequestCount: sess.RequestCount,
InputTokens: sess.InputTokens,
OutputTokens: sess.OutputTokens,
TotalTokens: sess.TotalTokens,
CachedInputTokens: sess.CachedInputTokens,
CacheCreationTokens: sess.CacheCreationTokens,
InputCostUsd: sess.InputCostUSD,
CachedInputCostUsd: sess.CachedInputCostUSD,
CacheCreationCostUsd: sess.CacheCreationCostUSD,
OutputCostUsd: sess.OutputCostUSD,
CostUsd: sess.TotalCostUSD(),
CacheCostUsd: sess.CacheCostUSD(),
Decision: sess.Decision,
Entries: entries,
StartedAt: sess.StartedAt,
EndedAt: sess.EndedAt,
RequestCount: sess.RequestCount,
InputTokens: sess.InputTokens,
OutputTokens: sess.OutputTokens,
TotalTokens: sess.TotalTokens,
CostUsd: sess.CostUSD,
Decision: sess.Decision,
Entries: entries,
}
out.SessionId = strPtr(sess.SessionID)
out.UserId = strPtr(sess.UserID)

View File

@@ -54,7 +54,7 @@ var accessLogSortFields = map[string]string{
"provider": "provider",
"status_code": "status_code",
"duration": "duration",
"cost_usd": CostUSDSQLExpr,
"cost_usd": "cost_usd",
"total_tokens": "total_tokens",
"user_id": "user_id",
"decision": "decision",
@@ -70,7 +70,7 @@ var accessLogSortFields = map[string]string{
var sessionSortExprs = map[string]string{ //nolint:gosec // G101 false positive: "total_tokens" sort key, not a credential
"timestamp": "MAX(timestamp)",
"started_at": "MIN(timestamp)",
"cost_usd": "SUM" + CostUSDSQLExpr,
"cost_usd": "SUM(cost_usd)",
"total_tokens": "SUM(total_tokens)",
"duration": "SUM(duration)",
"request_count": "COUNT(*)",

View File

@@ -1,124 +0,0 @@
package types
import (
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// costRow builds an access-log entry carrying only a cost breakdown — the rest
// of the row is irrelevant to the summation identities under test.
func costRow(id, session string, ts time.Time, in, cachedIn, cacheCreate, out float64) *AgentNetworkAccessLog {
return &AgentNetworkAccessLog{
ID: id,
SessionID: session,
Timestamp: ts,
InputCostUSD: in,
CachedInputCostUSD: cachedIn,
CacheCreationCostUSD: cacheCreate,
OutputCostUSD: out,
}
}
// TestAPIResponse_CostComponentsSumToAggregates is the contract a client adding
// up an API response depends on: within a single rendered object, the four
// per-bucket fields sum to cost_usd, and the two cache fields sum to
// cache_cost_usd. Uses rates that are not exactly representable in binary
// floating point, so the identity is checked against real arithmetic rather
// than round numbers.
func TestAPIResponse_CostComponentsSumToAggregates(t *testing.T) {
row := costRow("r1", "s1", time.Now(), 0.000768, 0.0002304, 0.00192, 0.003)
api := row.ToAPIResponse()
assert.InDelta(t, api.InputCostUsd+api.CachedInputCostUsd+api.CacheCreationCostUsd+api.OutputCostUsd,
api.CostUsd, 1e-12, "rendered buckets must sum to the rendered cost_usd")
assert.InDelta(t, api.CachedInputCostUsd+api.CacheCreationCostUsd, api.CacheCostUsd, 1e-12,
"rendered cache buckets must sum to the rendered cache_cost_usd")
assert.InDelta(t, 0.0059184, api.CostUsd, 1e-12, "total is the exact sum, not a separately rounded figure")
assert.InDelta(t, 0.0021504, api.CacheCostUsd, 1e-12, "cache cost is the exact sum of the two cache buckets")
}
// TestSessionSummary_SumsMatchSummedEntries proves a session summary equals the
// sum of the entries it renders: a client that adds up the entries itself must
// land on the same number the summary reports, per bucket and in total.
func TestSessionSummary_SumsMatchSummedEntries(t *testing.T) {
base := time.Date(2026, 5, 5, 10, 0, 0, 0, time.UTC)
entries := []*AgentNetworkAccessLog{
costRow("r1", "s1", base, 0.000768, 0.0002304, 0.00192, 0.003),
costRow("r2", "s1", base.Add(time.Minute), 0.000625, 0.0009375, 0, 0.005),
costRow("r3", "s1", base.Add(2*time.Minute), 0.0000016, 0, 0, 0.0000032),
}
sessions := FoldAccessLogSessions([]string{"s1"}, entries)
require.Len(t, sessions, 1)
sess := sessions[0].ToAPIResponse()
var wantInput, wantCachedInput, wantCacheCreation, wantOutput float64
for _, e := range entries {
wantInput += e.InputCostUSD
wantCachedInput += e.CachedInputCostUSD
wantCacheCreation += e.CacheCreationCostUSD
wantOutput += e.OutputCostUSD
}
assert.InDelta(t, wantInput, sess.InputCostUsd, 1e-12, "session input cost is the sum of its entries")
assert.InDelta(t, wantCachedInput, sess.CachedInputCostUsd, 1e-12, "session cache-read cost is the sum of its entries")
assert.InDelta(t, wantCacheCreation, sess.CacheCreationCostUsd, 1e-12, "session cache-write cost is the sum of its entries")
assert.InDelta(t, wantOutput, sess.OutputCostUsd, 1e-12, "session output cost is the sum of its entries")
assert.InDelta(t, wantInput+wantCachedInput+wantCacheCreation+wantOutput, sess.CostUsd, 1e-12,
"session total equals the summed entry buckets")
// Summing the rendered entries must give the same answer as reading the
// summary — the property a UI relies on when it totals a table itself.
var fromEntries float64
for _, e := range sess.Entries {
fromEntries += e.CostUsd
}
assert.InDelta(t, sess.CostUsd, fromEntries, 1e-12, "summary total must match the summed rendered entries")
// The sub-microdollar row must still contribute; it would vanish under
// 6-decimal quantisation.
assert.Greater(t, sess.InputCostUsd, 0.001393, "small-cost rows must not be quantised away")
}
// TestUsageBuckets_SumsMatchSummedRows proves the same identity one level up:
// a usage bucket equals the sum of the ledger rows folded into it, and the
// buckets together equal the whole range.
func TestUsageBuckets_SumsMatchSummedRows(t *testing.T) {
day1 := time.Date(2026, 5, 5, 9, 0, 0, 0, time.UTC)
day2 := time.Date(2026, 5, 6, 9, 0, 0, 0, time.UTC)
rows := []*AgentNetworkUsage{
{ID: "u1", Timestamp: day1, InputCostUSD: 0.000768, CachedInputCostUSD: 0.0002304, CacheCreationCostUSD: 0.00192, OutputCostUSD: 0.003},
{ID: "u2", Timestamp: day1.Add(time.Hour), InputCostUSD: 0.000625, CachedInputCostUSD: 0.0009375, OutputCostUSD: 0.005},
{ID: "u3", Timestamp: day2, InputCostUSD: 0.0000016, OutputCostUSD: 0.0000032},
}
buckets := AggregateUsageByGranularity(rows, UsageGranularityDay)
require.Len(t, buckets, 2, "two distinct days expected")
var total, cache float64
for _, b := range buckets {
api := b.ToAPIResponse()
assert.InDelta(t, api.InputCostUsd+api.CachedInputCostUsd+api.CacheCreationCostUsd+api.OutputCostUsd,
api.CostUsd, 1e-12, "each bucket's components must sum to its cost_usd")
total += api.CostUsd
cache += api.CacheCostUsd
}
var wantTotal, wantCache float64
for _, r := range rows {
wantTotal += r.TotalCostUSD()
wantCache += r.CacheCostUSD()
}
assert.InDelta(t, wantTotal, total, 1e-12, "buckets must sum to the total across all ledger rows")
assert.InDelta(t, wantCache, cache, 1e-12, "buckets must sum to the cache spend across all ledger rows")
// A month bucket over the same rows must total identically — regrouping
// changes the partition, never the sum.
monthly := AggregateUsageByGranularity(rows, UsageGranularityMonth)
require.Len(t, monthly, 1)
assert.InDelta(t, wantTotal, monthly[0].ToAPIResponse().CostUsd, 1e-12,
"re-bucketing at a different granularity must preserve the total")
}

View File

@@ -1,19 +0,0 @@
package types
// DiscoveredModel is a normalized model identifier returned by a provider's
// persisted upstream endpoint. Discovery does not persist models; the caller
// must explicitly update the provider to opt into the returned list.
type DiscoveredModel struct {
ID string
Label string
}
// ModelDiscoveryResult describes a successful proxy-executed discovery probe.
// RequestID correlates the management request with the proxy control message,
// and ProxyCluster identifies the network vantage point that ran it.
type ModelDiscoveryResult struct {
RequestID string
Source string
ProxyCluster string
Models []DiscoveredModel
}

View File

@@ -33,10 +33,6 @@ type Provider struct {
// the operator selected.
UpstreamURL string `gorm:"column:upstream_url"`
APIKey string `gorm:"column:api_key"`
// APIKeyProvided records whether api_key was present in the request. It is
// transient and lets updates distinguish omission (preserve) from an
// explicit empty value (clear an optional credential).
APIKeyProvided bool `gorm:"-" json:"-"`
// ExtraValues holds operator-typed values for catalog-declared
// ExtraHeaders (see catalog.Provider.ExtraHeaders). Keyed by
// header name (e.g. "x-portkey-config"); a non-empty value is
@@ -101,20 +97,15 @@ func NewProvider(accountID string) *Provider {
}
}
// FromAPIRequest applies the request payload onto the receiver. APIKeyProvided
// preserves the distinction between an omitted api_key and an explicit empty
// value; the manager applies the catalog-specific preserve/clear semantics.
// FromAPIRequest applies the request payload onto the receiver. The api_key
// is only overwritten when the caller provided one — empty/nil leaves the
// existing key intact, so updates can omit it.
func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) {
p.ProviderID = req.ProviderId
p.Name = req.Name
p.UpstreamURL = req.UpstreamUrl
p.APIKeyProvided = req.ApiKey != nil
if req.ApiKey != nil {
if strings.TrimSpace(*req.ApiKey) == "" {
p.APIKey = ""
} else {
p.APIKey = *req.ApiKey
}
if req.ApiKey != nil && strings.TrimSpace(*req.ApiKey) != "" {
p.APIKey = *req.ApiKey
}
if req.ExtraValues != nil {
// Replace the whole map (rather than merge) so unsetting a
@@ -187,7 +178,6 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
UpstreamUrl: p.UpstreamURL,
Models: models,
Enabled: p.Enabled,
HasApiKey: strings.TrimSpace(p.APIKey) != "",
SkipTlsVerification: p.SkipTLSVerification,
MetadataDisabled: p.MetadataDisabled,
CreatedAt: &created,

View File

@@ -77,37 +77,3 @@ func TestProvider_MetadataDisabled_RoundTrip(t *testing.T) {
assert.False(t, p.MetadataDisabled, "explicit false must clear metadata_disabled")
assert.False(t, p.ToAPIResponse().MetadataDisabled, "response must reflect the cleared value")
}
func TestProvider_APIKeyPresenceAndResponse(t *testing.T) {
base := func() *api.AgentNetworkProviderRequest {
return &api.AgentNetworkProviderRequest{
ProviderId: "ollama",
Name: "Ollama",
UpstreamUrl: "http://ollama.internal:11434",
}
}
p := NewProvider("acc-1")
p.FromAPIRequest(base())
assert.False(t, p.APIKeyProvided, "omission must remain distinguishable from an explicit clear")
assert.False(t, p.ToAPIResponse().HasApiKey)
key := "protected-endpoint-token"
withKey := base()
withKey.ApiKey = &key
p.FromAPIRequest(withKey)
assert.True(t, p.APIKeyProvided)
assert.Equal(t, key, p.APIKey)
assert.True(t, p.ToAPIResponse().HasApiKey)
empty := ""
clearKey := base()
clearKey.ApiKey = &empty
p.FromAPIRequest(clearKey)
assert.True(t, p.APIKeyProvided, "explicit empty must be carried to the manager as a clear")
assert.Empty(t, p.APIKey)
assert.False(t, p.ToAPIResponse().HasApiKey)
p.FromAPIRequest(base())
assert.False(t, p.APIKeyProvided, "a later omitted value must reset transient request presence")
}

View File

@@ -25,19 +25,8 @@ type AgentNetworkUsage struct {
InputTokens int64
OutputTokens int64
TotalTokens int64
// Prompt-cache buckets: read + write token counts.
CachedInputTokens int64
CacheCreationTokens int64
// Per-bucket cost breakdown, mirroring AgentNetworkAccessLog — the only
// cost state stored; total and cache portion are derived on read. Kept on
// the usage ledger too so spend can be attributed per bucket even for
// accounts with log collection turned off. See AgentNetworkAccessLog for
// why the columns carry a zero default.
InputCostUSD float64 `gorm:"not null;default:0"`
CachedInputCostUSD float64 `gorm:"not null;default:0"`
CacheCreationCostUSD float64 `gorm:"not null;default:0"`
OutputCostUSD float64 `gorm:"not null;default:0"`
CreatedAt time.Time
CostUSD float64
CreatedAt time.Time
}
// TableName keeps usage records in their own stripped table. Named
@@ -45,17 +34,6 @@ type AgentNetworkUsage struct {
// agent_network_usage table in a shared database.
func (AgentNetworkUsage) TableName() string { return "agent_network_request_usage" }
// TotalCostUSD is the request's total cost: the sum of the four per-bucket
// costs. Derived rather than stored so it cannot disagree with the breakdown.
func (u *AgentNetworkUsage) TotalCostUSD() float64 {
return u.InputCostUSD + u.CachedInputCostUSD + u.CacheCreationCostUSD + u.OutputCostUSD
}
// CacheCostUSD is the portion of the total billed for prompt-cache buckets.
func (u *AgentNetworkUsage) CacheCostUSD() float64 {
return u.CachedInputCostUSD + u.CacheCreationCostUSD
}
// AgentNetworkUsageGroup is the normalised many-to-many row linking a usage
// record to one authorising group, mirroring AgentNetworkAccessLogGroup so the
// usage overview can filter by group with a `group_id IN (...)` join.

View File

@@ -33,45 +33,21 @@ func ParseUsageGranularity(s string) UsageGranularity {
// AgentNetworkUsageBucket is one aggregated usage time bucket. PeriodStart is
// the UTC start of the bucket as YYYY-MM-DD.
type AgentNetworkUsageBucket struct {
PeriodStart string
InputTokens int64
OutputTokens int64
TotalTokens int64
CachedInputTokens int64
CacheCreationTokens int64
InputCostUSD float64
CachedInputCostUSD float64
CacheCreationCostUSD float64
OutputCostUSD float64
}
// TotalCostUSD is the bucket's total spend: the sum of the four per-bucket
// costs. Derived rather than accumulated separately so it cannot disagree with
// the components.
func (b *AgentNetworkUsageBucket) TotalCostUSD() float64 {
return b.InputCostUSD + b.CachedInputCostUSD + b.CacheCreationCostUSD + b.OutputCostUSD
}
// CacheCostUSD is the bucket's prompt-cache spend: cache reads plus writes.
func (b *AgentNetworkUsageBucket) CacheCostUSD() float64 {
return b.CachedInputCostUSD + b.CacheCreationCostUSD
PeriodStart string
InputTokens int64
OutputTokens int64
TotalTokens int64
CostUSD float64
}
// ToAPIResponse renders the bucket as the API representation.
func (b *AgentNetworkUsageBucket) ToAPIResponse() api.AgentNetworkUsageBucket {
return api.AgentNetworkUsageBucket{
PeriodStart: b.PeriodStart,
InputTokens: b.InputTokens,
OutputTokens: b.OutputTokens,
TotalTokens: b.TotalTokens,
CachedInputTokens: b.CachedInputTokens,
CacheCreationTokens: b.CacheCreationTokens,
InputCostUsd: b.InputCostUSD,
CachedInputCostUsd: b.CachedInputCostUSD,
CacheCreationCostUsd: b.CacheCreationCostUSD,
OutputCostUsd: b.OutputCostUSD,
CostUsd: b.TotalCostUSD(),
CacheCostUsd: b.CacheCostUSD(),
PeriodStart: b.PeriodStart,
InputTokens: b.InputTokens,
OutputTokens: b.OutputTokens,
TotalTokens: b.TotalTokens,
CostUsd: b.CostUSD,
}
}
@@ -108,12 +84,7 @@ func AggregateUsageByGranularity(rows []*AgentNetworkUsage, g UsageGranularity)
b.InputTokens += r.InputTokens
b.OutputTokens += r.OutputTokens
b.TotalTokens += r.TotalTokens
b.CachedInputTokens += r.CachedInputTokens
b.CacheCreationTokens += r.CacheCreationTokens
b.InputCostUSD += r.InputCostUSD
b.CachedInputCostUSD += r.CachedInputCostUSD
b.CacheCreationCostUSD += r.CacheCreationCostUSD
b.OutputCostUSD += r.OutputCostUSD
b.CostUSD += r.CostUSD
}
out := make([]*AgentNetworkUsageBucket, 0, len(byPeriod))

View File

@@ -4,16 +4,11 @@ package proxy
import (
"context"
"errors"
"time"
"github.com/netbirdio/netbird/shared/management/proto"
)
// ErrModelDiscoveryUnavailable is returned when no connected proxy in the
// requested cluster can execute a model-discovery request.
var ErrModelDiscoveryUnavailable = errors.New("model discovery proxy unavailable")
// Manager defines the interface for proxy operations
type Manager interface {
Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress string, accountID *string, capabilities *Capabilities) (*Proxy, error)
@@ -43,7 +38,6 @@ type OIDCValidationConfig struct {
// Controller is responsible for managing proxy clusters and routing service updates.
type Controller interface {
SendServiceUpdateToCluster(ctx context.Context, accountID string, update *proto.ProxyMapping, clusterAddr string)
DiscoverModels(ctx context.Context, accountID, clusterAddr string, req *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error)
GetOIDCValidationConfig() OIDCValidationConfig
RegisterProxyToCluster(ctx context.Context, clusterAddr, proxyID string) error
UnregisterProxyFromCluster(ctx context.Context, clusterAddr, proxyID string) error

View File

@@ -39,12 +39,6 @@ func (c *GRPCController) SendServiceUpdateToCluster(ctx context.Context, account
c.metrics.IncrementServiceUpdateSendCount(clusterAddr)
}
// DiscoverModels executes a correlated model-discovery request on one capable
// proxy connected to the requested cluster.
func (c *GRPCController) DiscoverModels(ctx context.Context, accountID, clusterAddr string, req *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
return c.proxyGRPCServer.DiscoverModels(ctx, accountID, clusterAddr, req)
}
// GetOIDCValidationConfig returns the OIDC validation configuration from the gRPC server.
func (c *GRPCController) GetOIDCValidationConfig() proxy.OIDCValidationConfig {
return c.proxyGRPCServer.GetOIDCValidationConfig()
@@ -56,12 +50,10 @@ func (c *GRPCController) RegisterProxyToCluster(ctx context.Context, clusterAddr
return nil
}
proxySet, _ := c.clusterProxies.LoadOrStore(clusterAddr, &sync.Map{})
_, alreadyRegistered := proxySet.(*sync.Map).LoadOrStore(proxyID, struct{}{})
proxySet.(*sync.Map).Store(proxyID, struct{}{})
log.WithContext(ctx).Debugf("Registered proxy %s to cluster %s", proxyID, clusterAddr)
if !alreadyRegistered {
c.metrics.IncrementProxyConnectionCount(clusterAddr)
}
c.metrics.IncrementProxyConnectionCount(clusterAddr)
return nil
}
@@ -72,10 +64,10 @@ func (c *GRPCController) UnregisterProxyFromCluster(ctx context.Context, cluster
return nil
}
if proxySet, ok := c.clusterProxies.Load(clusterAddr); ok {
if _, registered := proxySet.(*sync.Map).LoadAndDelete(proxyID); registered {
log.WithContext(ctx).Debugf("Unregistered proxy %s from cluster %s", proxyID, clusterAddr)
c.metrics.DecrementProxyConnectionCount(clusterAddr)
}
proxySet.(*sync.Map).Delete(proxyID)
log.WithContext(ctx).Debugf("Unregistered proxy %s from cluster %s", proxyID, clusterAddr)
c.metrics.DecrementProxyConnectionCount(clusterAddr)
}
return nil
}

View File

@@ -259,21 +259,6 @@ func (m *MockController) EXPECT() *MockControllerMockRecorder {
return m.recorder
}
// DiscoverModels mocks base method.
func (m *MockController) DiscoverModels(ctx context.Context, accountID, clusterAddr string, req *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "DiscoverModels", ctx, accountID, clusterAddr, req)
ret0, _ := ret[0].(*proto.ModelDiscoveryResult)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// DiscoverModels indicates an expected call of DiscoverModels.
func (mr *MockControllerMockRecorder) DiscoverModels(ctx, accountID, clusterAddr, req interface{}) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DiscoverModels", reflect.TypeOf((*MockController)(nil).DiscoverModels), ctx, accountID, clusterAddr, req)
}
// GetOIDCValidationConfig mocks base method.
func (m *MockController) GetOIDCValidationConfig() OIDCValidationConfig {
m.ctrl.T.Helper()

View File

@@ -5,6 +5,9 @@ import (
"strconv"
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/networkmap"
@@ -163,7 +166,7 @@ func (e *componentEncoder) indexAllPeers() {
}
}
func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
func (e *componentEncoder) appendPeer(p *nbpeer.Peer) uint32 {
if idx, ok := e.peerOrder[p.ID]; ok {
return idx
}
@@ -177,7 +180,7 @@ func (e *componentEncoder) appendPeer(p *types.ComponentPeer) uint32 {
// (c.RouterPeers may contain peers not in c.Peers when validation rules drop
// them) and returns their wire indexes for the RouterPeerIndexes field. Must
// run before any encoder that resolves peer ids via e.peerOrder.
func (e *componentEncoder) indexRouterPeers(routers map[string]*types.ComponentPeer) []uint32 {
func (e *componentEncoder) indexRouterPeers(routers map[string]*nbpeer.Peer) []uint32 {
if len(routers) == 0 {
return nil
}
@@ -511,7 +514,7 @@ func encodeCustomZones(zones []nbdns.CustomZone) []*proto.CustomZone {
return out
}
func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentResource) []*proto.NetworkResourceRaw {
func (e *componentEncoder) encodeNetworkResources(resources []*resourceTypes.NetworkResource) []*proto.NetworkResourceRaw {
if len(resources) == 0 {
return nil
}
@@ -540,7 +543,7 @@ func (e *componentEncoder) encodeNetworkResources(resources []*types.ComponentRe
return out
}
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*types.ComponentRouter) map[string]*proto.NetworkRouterList {
func (e *componentEncoder) encodeRoutersMap(routersMap map[string]map[string]*routerTypes.NetworkRouter) map[string]*proto.NetworkRouterList {
if len(routersMap) == 0 {
return nil
}
@@ -689,20 +692,20 @@ func toAccountNetwork(n *types.Network) *proto.AccountNetwork {
return out
}
func toPeerCompact(p *types.ComponentPeer) *proto.PeerCompact {
func toPeerCompact(p *nbpeer.Peer) *proto.PeerCompact {
pc := &proto.PeerCompact{
WgPubKey: decodeWgKey(p.Key),
SshPubKey: []byte(p.SSHKey),
DnsLabel: p.DNSLabel,
AgentVersion: p.AgentVersion,
AddedWithSsoLogin: p.AddedWithSSOLogin,
AgentVersion: p.Meta.WtVersion,
AddedWithSsoLogin: p.UserID != "",
LoginExpirationEnabled: p.LoginExpirationEnabled,
SshEnabled: p.SSHEnabled,
SupportsIpv6: p.SupportsIPv6,
SupportsSourcePrefixes: p.SupportsSourcePrefixes,
ServerSshAllowed: p.ServerSSHAllowed,
SupportsIpv6: p.SupportsIPv6(),
SupportsSourcePrefixes: p.SupportsSourcePrefixes(),
ServerSshAllowed: p.Meta.Flags.ServerSSHAllowed,
}
if !p.LastLogin.IsZero() {
if p.LastLogin != nil {
pc.LastLoginUnixNano = p.LastLogin.UnixNano()
}
switch {

View File

@@ -15,6 +15,9 @@ import (
goproto "google.golang.org/protobuf/proto"
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/proto"
@@ -152,28 +155,29 @@ func envelopesEquivalent(a, b *proto.NetworkMapEnvelope) bool {
}
func newTestComponents() *types.NetworkMapComponents {
peerA := &types.ComponentPeer{
ID: "peer-a",
Key: testWgKeyA,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peera",
SSHKey: "ssh-a",
AgentVersion: "0.40.0",
peerA := &nbpeer.Peer{
ID: "peer-a",
Key: testWgKeyA,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peera",
SSHKey: "ssh-a",
Status: &nbpeer.PeerStatus{Connected: true, LastSeen: time.Now()},
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
peerB := &types.ComponentPeer{
ID: "peer-b",
Key: testWgKeyB,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
DNSLabel: "peerb",
AgentVersion: "0.25.0",
peerB := &nbpeer.Peer{
ID: "peer-b",
Key: testWgKeyB,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2}),
DNSLabel: "peerb",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.25.0"},
}
peerC := &types.ComponentPeer{
ID: "peer-c",
Key: testWgKeyC,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc",
AgentVersion: "0.40.0",
peerC := &nbpeer.Peer{
ID: "peer-c",
Key: testWgKeyC,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
return &types.NetworkMapComponents{
@@ -187,12 +191,12 @@ func newTestComponents() *types.NetworkMapComponents {
PeerLoginExpirationEnabled: true,
PeerLoginExpiration: 2 * time.Hour,
},
Peers: map[string]*types.ComponentPeer{
Peers: map[string]*nbpeer.Peer{
"peer-a": peerA,
"peer-b": peerB,
"peer-c": peerC,
},
Groups: map[string]*types.ComponentGroup{
Groups: map[string]*types.Group{
"group-src": {ID: "group-src", PublicID: "1", Name: "Src", Peers: []string{"peer-a"}},
"group-dst": {ID: "group-dst", PublicID: "2", Name: "Dst", Peers: []string{"peer-b", "peer-c"}},
},
@@ -211,7 +215,7 @@ func newTestComponents() *types.NetworkMapComponents {
}},
},
},
RouterPeers: map[string]*types.ComponentPeer{"peer-c": peerC},
RouterPeers: map[string]*nbpeer.Peer{"peer-c": peerC},
}
}
@@ -377,12 +381,12 @@ func TestEncodeNetworkMapEnvelope_MalformedWgKey(t *testing.T) {
func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
c := newTestComponents()
v6Only := &types.ComponentPeer{
ID: "peer-v6",
Key: testWgKeyA,
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
DNSLabel: "peerv6",
AgentVersion: "0.40.0",
v6Only := &nbpeer.Peer{
ID: "peer-v6",
Key: testWgKeyA,
IPv6: netip.AddrFrom16([16]byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 9}),
DNSLabel: "peerv6",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.Peers["peer-v6"] = v6Only
@@ -401,11 +405,11 @@ func TestEncodeNetworkMapEnvelope_IPv6OnlyPeer(t *testing.T) {
func TestEncodeNetworkMapEnvelope_PeerWithoutIP(t *testing.T) {
c := newTestComponents()
c.Peers["peer-noip"] = &types.ComponentPeer{
ID: "peer-noip",
Key: testWgKeyA,
DNSLabel: "peernoip",
AgentVersion: "0.40.0",
c.Peers["peer-noip"] = &nbpeer.Peer{
ID: "peer-noip",
Key: testWgKeyA,
DNSLabel: "peernoip",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -440,9 +444,9 @@ func TestEncodeNetworkMapEnvelope_EmptyInput(t *testing.T) {
func TestEncodeNetworkMapEnvelope_PeerLoginExpirationFields(t *testing.T) {
c := newTestComponents()
now := time.Date(2024, 1, 2, 3, 4, 5, 0, time.UTC)
c.Peers["peer-a"].AddedWithSSOLogin = true
c.Peers["peer-a"].UserID = "user-1"
c.Peers["peer-a"].LoginExpirationEnabled = true
c.Peers["peer-a"].LastLogin = now
c.Peers["peer-a"].LastLogin = &now
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -553,7 +557,7 @@ func TestEncodeNetworkMapEnvelope_ResourceOnlyPolicyShippedAndIndexed(t *testing
}
// Resource must appear in components.NetworkResources with a seq id —
// encoder uses that to translate the xid map key to uint32.
c.NetworkResources = []*types.ComponentResource{
c.NetworkResources = []*resourceTypes.NetworkResource{
{ID: "resource-x", PublicID: "77", Name: "res-x", Enabled: true},
}
@@ -621,11 +625,11 @@ func TestEncodeNetworkMapEnvelope_PostureFailedPeers(t *testing.T) {
func TestEncodeNetworkMapEnvelope_RoutersMap(t *testing.T) {
c := newTestComponents()
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
c.RoutersMap = map[string]map[string]*routerTypes.NetworkRouter{
"net-1": {
"peer-c": {
PublicID: "200",
Peer: "peer-c", Masquerade: true, Metric: 10, Enabled: true,
ID: "router-1", PublicID: "200",
Peer: "peer-c", Masquerade: true, Metric: 10, Enabled: true,
},
},
}
@@ -651,14 +655,14 @@ func TestEncodeNetworkMapEnvelope_RouterPeerNotInComponentsPeers(t *testing.T) {
// peer_index reference must still resolve.
c := newTestComponents()
delete(c.Peers, "peer-c")
routerPeer := &types.ComponentPeer{
routerPeer := &nbpeer.Peer{
ID: "peer-c", Key: testWgKeyC, IP: netip.AddrFrom4([4]byte{100, 64, 0, 3}),
DNSLabel: "peerc", AgentVersion: "0.40.0",
DNSLabel: "peerc", Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
c.RouterPeers = map[string]*types.ComponentPeer{"peer-c": routerPeer}
c.RouterPeers = map[string]*nbpeer.Peer{"peer-c": routerPeer}
c.NetworkXIDToPublicID = map[string]string{"net-1": "5"}
c.RoutersMap = map[string]map[string]*types.ComponentRouter{
"net-1": {"peer-c": {PublicID: "1", Peer: "peer-c", Enabled: true}},
c.RoutersMap = map[string]map[string]*routerTypes.NetworkRouter{
"net-1": {"peer-c": {ID: "r-1", PublicID: "1", Peer: "peer-c", Enabled: true}},
}
full := EncodeNetworkMapEnvelope(ComponentsEnvelopeInput{Components: c}).GetFull()
@@ -691,9 +695,9 @@ func TestToProxyPatch_EmptyInputReturnsNil(t *testing.T) {
func TestToProxyPatch_PopulatesAllFields(t *testing.T) {
nm := &types.NetworkMap{
Peers: []*types.ComponentPeer{{
Peers: []*nbpeer.Peer{{
ID: "ext-peer", Key: testWgKeyA, IP: netip.AddrFrom4([4]byte{100, 64, 0, 9}),
DNSLabel: "extpeer", AgentVersion: "0.40.0",
DNSLabel: "extpeer", Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}},
FirewallRules: []*types.FirewallRule{{
PeerIP: "100.64.0.9", Action: "accept", Direction: 0, Protocol: "tcp",
@@ -776,6 +780,6 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
func emptyNetworkMapComponents() *types.NetworkMapComponents {
return types.EmptyNetworkMapComponents(
&types.NetworkMapComponents{
PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}},
PeerID: "peer-id", Peers: map[string]*nbpeer.Peer{"peer-id": {}}},
)
}

View File

@@ -19,10 +19,10 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
mkComponents := func(rule *types.PolicyRule, sshEnabled bool) (*types.NetworkMapComponents, *nbpeer.Peer) {
peer := &nbpeer.Peer{ID: targetPeerID, SSHEnabled: sshEnabled}
group := &types.ComponentGroup{ID: targetGroupID, Name: "dst", Peers: []string{targetPeerID}}
group := &types.Group{ID: targetGroupID, Name: "dst", Peers: []string{targetPeerID}}
return &types.NetworkMapComponents{
Peers: map[string]*types.ComponentPeer{targetPeerID: peer.ToComponent()},
Groups: map[string]*types.ComponentGroup{targetGroupID: group},
Peers: map[string]*nbpeer.Peer{targetPeerID: peer},
Groups: map[string]*types.Group{targetGroupID: group},
Policies: []*types.Policy{{
ID: "p",
Enabled: true,
@@ -158,8 +158,8 @@ func TestComputeSSHEnabledForPeer(t *testing.T) {
func TestComputeSSHEnabledForPeer_TargetMissingFromComponents(t *testing.T) {
peer := &nbpeer.Peer{ID: "missing", SSHEnabled: true}
c := &types.NetworkMapComponents{
Peers: map[string]*types.ComponentPeer{}, // target peer NOT present
Groups: map[string]*types.ComponentGroup{
Peers: map[string]*nbpeer.Peer{}, // target peer NOT present
Groups: map[string]*types.Group{
"g": {ID: "g", Peers: []string{"missing"}},
},
Policies: []*types.Policy{{

View File

@@ -1,404 +0,0 @@
package grpc
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
rpproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
"github.com/netbirdio/netbird/shared/management/proto"
nbstatus "github.com/netbirdio/netbird/shared/management/status"
)
type modelDiscoveryCall struct {
result *proto.ModelDiscoveryResult
err error
}
func newModelDiscoveryTestServer(t *testing.T) (*ProxyServiceServer, *testProxyController) {
t.Helper()
controller := newTestProxyController()
server := &ProxyServiceServer{
proxyController: controller,
modelDiscoveryPending: make(map[string]*pendingModelDiscovery),
}
return server, controller
}
func addModelDiscoveryTestConnection(
t *testing.T,
server *ProxyServiceServer,
controller *testProxyController,
proxyID, sessionID, cluster string,
accountID *string,
supportsDiscovery bool,
) (*proxyConnection, context.CancelFunc) {
t.Helper()
ctx, cancel := context.WithCancel(context.Background())
ready := make(chan struct{})
close(ready)
conn := &proxyConnection{
proxyID: proxyID,
sessionID: sessionID,
address: cluster,
accountID: accountID,
capabilities: &proto.ProxyCapabilities{
SupportsModelDiscovery: &supportsDiscovery,
},
syncStream: &syncRecordingStream{},
modelDiscoveryChan: make(chan *proto.ModelDiscoveryRequest, modelDiscoveryQueueSize),
ready: ready,
ctx: ctx,
cancel: cancel,
}
server.connectedProxies.Store(proxyID, conn)
require.NoError(t, controller.RegisterProxyToCluster(context.Background(), cluster, proxyID))
t.Cleanup(func() {
cancel()
server.connectedProxies.Delete(proxyID)
_ = controller.UnregisterProxyFromCluster(context.Background(), cluster, proxyID)
})
return conn, cancel
}
func callDiscoverModels(
server *ProxyServiceServer,
ctx context.Context,
accountID, cluster string,
req *proto.ModelDiscoveryRequest,
) <-chan modelDiscoveryCall {
done := make(chan modelDiscoveryCall, 1)
go func() {
result, err := server.DiscoverModels(ctx, accountID, cluster, req)
done <- modelDiscoveryCall{result: result, err: err}
}()
return done
}
func TestDiscoverModels_CorrelatesResultAndCopiesSensitiveRequest(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
callerReq := &proto.ModelDiscoveryRequest{
RequestId: "caller-controlled",
UpstreamUrl: "http://ollama.internal:11434",
AuthHeaderName: "Authorization",
AuthHeaderValue: "Bearer secret",
SkipTlsVerify: true,
OllamaFallback: true,
}
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", callerReq)
wireReq := <-conn.modelDiscoveryChan
assert.NotEmpty(t, wireReq.GetRequestId())
assert.NotEqual(t, callerReq.GetRequestId(), wireReq.GetRequestId(), "management must own correlation IDs")
assert.Equal(t, callerReq.GetUpstreamUrl(), wireReq.GetUpstreamUrl())
assert.Equal(t, callerReq.GetAuthHeaderName(), wireReq.GetAuthHeaderName())
assert.Equal(t, callerReq.GetAuthHeaderValue(), wireReq.GetAuthHeaderValue())
assert.Equal(t, callerReq.GetSkipTlsVerify(), wireReq.GetSkipTlsVerify())
assert.Equal(t, callerReq.GetOllamaFallback(), wireReq.GetOllamaFallback())
assert.Equal(t, "caller-controlled", callerReq.GetRequestId(), "caller request must not be mutated")
want := &proto.ModelDiscoveryResult{
RequestId: wireReq.GetRequestId(),
Models: []*proto.ModelDiscoveryModel{
{Id: "llama3.2:latest", Label: "llama3.2:latest"},
},
Source: "openai_v1_models",
}
server.completeModelDiscovery(conn, want)
got := <-done
require.NoError(t, got.err)
assert.Equal(t, want, got.result)
}
func TestDiscoverModels_IgnoresMismatchedCorrelation(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
wireReq := <-conn.modelDiscoveryChan
server.completeModelDiscovery(conn, &proto.ModelDiscoveryResult{RequestId: "wrong-request"})
select {
case got := <-done:
t.Fatalf("mismatched request completed discovery: %+v", got)
default:
}
wrongSession := *conn
wrongSession.sessionID = "session-b"
server.completeModelDiscovery(&wrongSession, &proto.ModelDiscoveryResult{RequestId: wireReq.GetRequestId()})
select {
case got := <-done:
t.Fatalf("mismatched proxy session completed discovery: %+v", got)
default:
}
server.completeModelDiscovery(conn, &proto.ModelDiscoveryResult{RequestId: wireReq.GetRequestId()})
got := <-done
require.NoError(t, got.err)
require.NotNil(t, got.result)
}
func TestDiscoverModels_ReturnsProxyReportedErrorForManagerContext(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
wireReq := <-conn.modelDiscoveryChan
server.completeModelDiscovery(conn, &proto.ModelDiscoveryResult{
RequestId: wireReq.GetRequestId(),
Error: "upstream returned HTTP 401",
})
got := <-done
require.NoError(t, got.err)
require.NotNil(t, got.result)
assert.Equal(t, "upstream returned HTTP 401", got.result.GetError())
assert.Equal(t, wireReq.GetRequestId(), got.result.GetRequestId())
}
func TestDiscoverModels_FiltersCapabilityAndAccountScope(t *testing.T) {
tests := []struct {
name string
connectionAccount *string
requestAccount string
supported bool
}{
{
name: "capability absent",
requestAccount: "account-a",
supported: false,
},
{
name: "BYOP account mismatch",
connectionAccount: stringPointer("account-b"),
requestAccount: "account-a",
supported: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com",
tt.connectionAccount, tt.supported,
)
result, err := server.DiscoverModels(
context.Background(),
tt.requestAccount,
"cluster.example.com",
&proto.ModelDiscoveryRequest{UpstreamUrl: "http://ollama.internal:11434"},
)
require.Nil(t, result)
requireStatusType(t, err, nbstatus.PreconditionFailed)
})
}
}
func TestDiscoverModels_RejectsStaleClusterMembershipAfterReconnect(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "new-session", "new-cluster.example.com", nil, true,
)
// Simulate the stale membership left behind when this proxy ID previously
// connected to another cluster and that superseded session skipped cleanup.
require.NoError(t, controller.RegisterProxyToCluster(
context.Background(),
"old-cluster.example.com",
"proxy-a",
))
result, err := server.DiscoverModels(
context.Background(),
"account-a",
"old-cluster.example.com",
&proto.ModelDiscoveryRequest{UpstreamUrl: "http://ollama.internal:11434"},
)
require.Nil(t, result)
requireStatusType(t, err, nbstatus.PreconditionFailed)
select {
case request := <-conn.modelDiscoveryChan:
t.Fatalf("sent credentials to connection in a different cluster: %+v", request)
default:
}
}
func TestDiscoverModels_CancelAndDisconnectCleanPendingRequest(t *testing.T) {
t.Run("caller cancellation", func(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
ctx, cancel := context.WithCancel(context.Background())
done := callDiscoverModels(server, ctx, "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
<-conn.modelDiscoveryChan
cancel()
got := <-done
require.Nil(t, got.result)
requireStatusType(t, got.err, nbstatus.PreconditionFailed)
assert.Empty(t, server.modelDiscoveryPending)
})
t.Run("proxy disconnect", func(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
<-conn.modelDiscoveryChan
server.failModelDiscoveries(conn, rpproxy.ErrModelDiscoveryUnavailable)
got := <-done
require.Nil(t, got.result)
requireStatusType(t, got.err, nbstatus.PreconditionFailed)
assert.Empty(t, server.modelDiscoveryPending)
})
}
func TestDiscoverModels_UsesNextCapableProxyWhenFirstQueueIsFull(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
first, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
second, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-b", "session-b", "cluster.example.com", nil, true,
)
first.modelDiscoveryChan = make(chan *proto.ModelDiscoveryRequest, 1)
first.modelDiscoveryChan <- &proto.ModelDiscoveryRequest{RequestId: "occupied"}
done := callDiscoverModels(server, context.Background(), "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
wireReq := <-second.modelDiscoveryChan
server.completeModelDiscovery(second, &proto.ModelDiscoveryResult{RequestId: wireReq.GetRequestId()})
got := <-done
require.NoError(t, got.err)
require.NotNil(t, got.result)
}
func TestSenderSkipsCanceledQueuedModelDiscovery(t *testing.T) {
server, controller := newModelDiscoveryTestServer(t)
conn, _ := addModelDiscoveryTestConnection(
t, server, controller, "proxy-a", "session-a", "cluster.example.com", nil, true,
)
stream := conn.syncStream.(*syncRecordingStream)
ctx, cancel := context.WithCancel(context.Background())
done := callDiscoverModels(server, ctx, "account-a", "cluster.example.com", &proto.ModelDiscoveryRequest{
UpstreamUrl: "http://ollama.internal:11434",
})
require.Eventually(t, func() bool {
return len(conn.modelDiscoveryChan) == 1
}, time.Second, time.Millisecond)
cancel()
got := <-done
require.Nil(t, got.result)
requireStatusType(t, got.err, nbstatus.PreconditionFailed)
liveRequest := &proto.ModelDiscoveryRequest{RequestId: "live-request"}
livePending := &pendingModelDiscovery{
proxyID: conn.proxyID,
sessionID: conn.sessionID,
done: make(chan modelDiscoveryCompletion, 1),
}
server.registerModelDiscovery(liveRequest.GetRequestId(), livePending)
defer server.removeModelDiscovery(liveRequest.GetRequestId(), livePending)
conn.modelDiscoveryChan <- liveRequest
errCh := make(chan error, 1)
go server.sender(conn, errCh)
require.Eventually(t, func() bool {
stream.mu.Lock()
defer stream.mu.Unlock()
return len(stream.sent) == 1
}, time.Second, time.Millisecond)
stream.mu.Lock()
sent := append([]*proto.SyncMappingsResponse(nil), stream.sent...)
stream.mu.Unlock()
require.Len(t, sent, 1)
assert.Equal(t, "live-request", sent[0].GetModelDiscoveryRequest().GetRequestId())
}
func TestDrainRecv_DispatchesModelDiscoveryResult(t *testing.T) {
server, _ := newModelDiscoveryTestServer(t)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
conn := &proxyConnection{
proxyID: "proxy-a",
sessionID: "session-a",
ctx: ctx,
cancel: cancel,
}
pending := &pendingModelDiscovery{
proxyID: conn.proxyID,
sessionID: conn.sessionID,
done: make(chan modelDiscoveryCompletion, 1),
}
server.registerModelDiscovery("request-a", pending)
stream := &syncRecordingStream{
recvMsgs: []*proto.SyncMappingsRequest{
{
Msg: &proto.SyncMappingsRequest_ModelDiscoveryResult{
ModelDiscoveryResult: &proto.ModelDiscoveryResult{RequestId: "request-a"},
},
},
},
}
errCh := make(chan error, 1)
go server.drainRecv(conn, stream, errCh)
completion := <-pending.done
require.NoError(t, completion.err)
require.NotNil(t, completion.result)
assert.Equal(t, "request-a", completion.result.GetRequestId())
}
func TestSendModelDiscoveryRequest_UsesOutOfBandSyncField(t *testing.T) {
stream := &syncRecordingStream{}
conn := &proxyConnection{syncStream: stream}
req := &proto.ModelDiscoveryRequest{RequestId: "request-a"}
require.NoError(t, conn.sendModelDiscoveryRequest(req))
require.Len(t, stream.sent, 1)
assert.Empty(t, stream.sent[0].GetMapping())
assert.Equal(t, req, stream.sent[0].GetModelDiscoveryRequest())
}
func stringPointer(value string) *string {
return &value
}
func requireStatusType(t *testing.T, err error, want nbstatus.Type) {
t.Helper()
require.Error(t, err)
statusErr, ok := nbstatus.FromError(err)
require.True(t, ok)
assert.Equal(t, want, statusErr.Type())
}

View File

@@ -15,7 +15,6 @@ import (
"net/http"
"net/url"
"os"
"sort"
"strconv"
"strings"
"sync"
@@ -30,7 +29,6 @@ import (
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/peers"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
@@ -38,6 +36,7 @@ import (
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
"github.com/netbirdio/netbird/management/server/idp"
"github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/management/server/users"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
@@ -85,10 +84,6 @@ type ProxyServiceServer struct {
// Map of connected proxies: proxy_id -> proxy connection
connectedProxies sync.Map
// modelDiscoveryPending correlates out-of-band model-discovery results
// received on SyncMappings with the management request waiting for them.
modelDiscoveryMu sync.Mutex
modelDiscoveryPending map[string]*pendingModelDiscovery
// Manager for access logs
accessLogManager accesslogs.Manager
@@ -148,8 +143,6 @@ const defaultProxyTokenTTL = 5 * time.Minute
const defaultSnapshotBatchSize = 500
const modelDiscoveryQueueSize = 16
func snapshotBatchSizeFromEnv() int {
if v := os.Getenv("NB_PROXY_SNAPSHOT_BATCH_SIZE"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
@@ -178,26 +171,10 @@ type proxyConnection struct {
stream proto.ProxyService_GetMappingUpdateServer
// syncStream is set when the proxy connected via SyncMappings.
// When non-nil, the sender goroutine uses this instead of stream.
syncStream proto.ProxyService_SyncMappingsServer
sendChan chan *proto.GetMappingUpdateResponse
modelDiscoveryChan chan *proto.ModelDiscoveryRequest
// ready closes after the initial snapshot completes and the sender
// goroutine is about to start. Discovery never competes with snapshot
// delivery on the stream.
ready chan struct{}
ctx context.Context
cancel context.CancelFunc
}
type modelDiscoveryCompletion struct {
result *proto.ModelDiscoveryResult
err error
}
type pendingModelDiscovery struct {
proxyID string
sessionID string
done chan modelDiscoveryCompletion
syncStream proto.ProxyService_SyncMappingsServer
sendChan chan *proto.GetMappingUpdateResponse
ctx context.Context
cancel context.CancelFunc
}
func enforceAccountScope(ctx context.Context, requestAccountID string) error {
@@ -215,18 +192,17 @@ func enforceAccountScope(ctx context.Context, requestAccountID string) error {
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, pkceStore *PKCEVerifierStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
ctx, cancel := context.WithCancel(context.Background())
s := &ProxyServiceServer{
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
proxyManager: proxyMgr,
tokenChecker: tokenChecker,
snapshotBatchSize: snapshotBatchSizeFromEnv(),
modelDiscoveryPending: make(map[string]*pendingModelDiscovery),
cancel: cancel,
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
proxyManager: proxyMgr,
tokenChecker: tokenChecker,
snapshotBatchSize: snapshotBatchSizeFromEnv(),
cancel: cancel,
}
go s.cleanupStaleProxies(ctx)
return s
@@ -416,7 +392,6 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest
return fmt.Errorf("send snapshot to proxy %s: %w", params.proxyID, err)
}
close(conn.ready)
errChan := make(chan error, 2)
go s.sender(conn, errChan)
@@ -450,10 +425,9 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings
return fmt.Errorf("send snapshot to proxy %s: %w", params.proxyID, err)
}
close(conn.ready)
errChan := make(chan error, 2)
go s.sender(conn, errChan)
go s.drainRecv(conn, stream, errChan)
go s.drainRecv(stream, errChan)
return s.serveProxyConnection(conn, proxyRecord, errChan, true)
}
@@ -522,8 +496,6 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
connSeed.tokenID = tokenID
connSeed.capabilities = params.capabilities
connSeed.sendChan = make(chan *proto.GetMappingUpdateResponse, 100)
connSeed.modelDiscoveryChan = make(chan *proto.ModelDiscoveryRequest, modelDiscoveryQueueSize)
connSeed.ready = make(chan struct{})
connSeed.ctx = connCtx
connSeed.cancel = cancel
@@ -571,13 +543,10 @@ func (s *ProxyServiceServer) supersedePriorConnection(proxyID, newSessionID stri
// cleanupFailedSnapshot removes the connection from the cluster and store
// after a snapshot send failure.
func (s *ProxyServiceServer) cleanupFailedSnapshot(ctx context.Context, conn *proxyConnection) {
s.failModelDiscoveries(conn, proxy.ErrModelDiscoveryUnavailable)
if s.connectedProxies.CompareAndDelete(conn.proxyID, conn) {
if err := s.proxyController.UnregisterProxyFromCluster(context.Background(), conn.address, conn.proxyID); err != nil {
log.WithContext(ctx).Debugf("cleanup after snapshot failure for proxy %s: %v", conn.proxyID, err)
}
} else {
s.unregisterSupersededProxyCluster(ctx, conn)
}
conn.cancel()
if err := s.proxyManager.Disconnect(context.Background(), conn.proxyID, conn.sessionID); err != nil {
@@ -585,19 +554,15 @@ func (s *ProxyServiceServer) cleanupFailedSnapshot(ctx context.Context, conn *pr
}
}
// drainRecv consumes post-snapshot messages from a bidirectional stream.
// Incremental mapping acks need no action; model-discovery results are
// dispatched to the correlated management caller.
func (s *ProxyServiceServer) drainRecv(conn *proxyConnection, stream proto.ProxyService_SyncMappingsServer, errChan chan<- error) {
// drainRecv consumes and discards messages from a bidirectional stream.
// The proxy sends an ack for every incremental update; we don't need them
// after the snapshot phase. Recv errors are forwarded to errChan.
func (s *ProxyServiceServer) drainRecv(stream proto.ProxyService_SyncMappingsServer, errChan chan<- error) {
for {
msg, err := stream.Recv()
if err != nil {
if _, err := stream.Recv(); err != nil {
errChan <- err
return
}
if result := msg.GetModelDiscoveryResult(); result != nil {
s.completeModelDiscovery(conn, result)
}
}
}
@@ -634,9 +599,7 @@ func (s *ProxyServiceServer) serveProxyConnection(conn *proxyConnection, proxyRe
// disconnectProxy removes the connection from cluster and store, unless it
// has already been superseded by a newer connection.
func (s *ProxyServiceServer) disconnectProxy(conn *proxyConnection) {
s.failModelDiscoveries(conn, proxy.ErrModelDiscoveryUnavailable)
if !s.connectedProxies.CompareAndDelete(conn.proxyID, conn) {
s.unregisterSupersededProxyCluster(context.Background(), conn)
log.Infof("Proxy %s session %s: skipping cleanup, superseded by new connection", conn.proxyID, conn.sessionID)
conn.cancel()
return
@@ -653,26 +616,6 @@ func (s *ProxyServiceServer) disconnectProxy(conn *proxyConnection) {
log.Infof("Proxy %s session %s disconnected", conn.proxyID, conn.sessionID)
}
// unregisterSupersededProxyCluster removes membership owned only by an old
// session after the same proxy ID reconnects at a different address. When the
// address is unchanged, the membership is shared with the replacement and
// must remain registered.
func (s *ProxyServiceServer) unregisterSupersededProxyCluster(ctx context.Context, conn *proxyConnection) {
if conn == nil || s.proxyController == nil {
return
}
currentValue, ok := s.connectedProxies.Load(conn.proxyID)
if ok {
current := currentValue.(*proxyConnection)
if current == conn || current.address == conn.address {
return
}
}
if err := s.proxyController.UnregisterProxyFromCluster(ctx, conn.address, conn.proxyID); err != nil {
log.WithContext(ctx).Debugf("cleanup superseded cluster membership for proxy %s: %v", conn.proxyID, err)
}
}
// sendSnapshotSync sends the initial snapshot with back-pressure: it sends
// one batch, then waits for the proxy to ack before sending the next.
func (s *ProxyServiceServer) sendSnapshotSync(ctx context.Context, conn *proxyConnection, stream proto.ProxyService_SyncMappingsServer) error {
@@ -909,20 +852,6 @@ func (s *ProxyServiceServer) sender(conn *proxyConnection, errChan chan<- error)
return
}
log.WithContext(conn.ctx).Tracef("Send response to proxy %s", conn.proxyID)
case req := <-conn.modelDiscoveryChan:
// The API caller may have timed out while this request waited
// behind other stream work. Pending correlation is removed on
// cancellation, so skip stale probes before they reach the
// proxy and expose a stored credential unnecessarily.
if !s.isModelDiscoveryPendingFor(conn, req.GetRequestId()) {
continue
}
if err := conn.sendModelDiscoveryRequest(req); err != nil {
errChan <- err
log.WithContext(conn.ctx).Tracef("Failed to send model discovery request to proxy %s: %v", conn.proxyID, err)
return
}
log.WithContext(conn.ctx).Tracef("Sent model discovery request to proxy %s", conn.proxyID)
case <-conn.ctx.Done():
return
}
@@ -940,15 +869,6 @@ func (conn *proxyConnection) sendResponse(resp *proto.GetMappingUpdateResponse)
return conn.stream.Send(resp)
}
func (conn *proxyConnection) sendModelDiscoveryRequest(req *proto.ModelDiscoveryRequest) error {
if conn.syncStream == nil {
return proxy.ErrModelDiscoveryUnavailable
}
return conn.syncStream.Send(&proto.SyncMappingsResponse{
ModelDiscoveryRequest: req,
})
}
// SendAccessLog processes access log from proxy
func (s *ProxyServiceServer) SendAccessLog(ctx context.Context, req *proto.SendAccessLogRequest) (*proto.SendAccessLogResponse, error) {
accessLog := req.GetLog()
@@ -1097,181 +1017,6 @@ func (s *ProxyServiceServer) GetConnectedProxyURLs() []string {
return urls
}
// DiscoverModels sends one correlated model-discovery request to a capable
// proxy in clusterAddr and waits for its result. Only the stored connection's
// cluster/account scope is used for routing; no routing identity is carried in
// the request delivered to the proxy.
func (s *ProxyServiceServer) DiscoverModels(ctx context.Context, accountID, clusterAddr string, req *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
if req == nil {
return nil, nbstatus.Errorf(nbstatus.InvalidArgument, "model discovery request is required")
}
if strings.TrimSpace(accountID) == "" {
return nil, nbstatus.Errorf(nbstatus.InvalidArgument, "account ID is required for model discovery")
}
if strings.TrimSpace(clusterAddr) == "" {
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "agent network proxy cluster is not configured")
}
if s.proxyController == nil {
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery is unavailable for the configured proxy cluster")
}
proxyIDs := s.proxyController.GetProxiesForCluster(clusterAddr)
sort.Strings(proxyIDs)
request := &proto.ModelDiscoveryRequest{
RequestId: uuid.NewString(),
UpstreamUrl: req.GetUpstreamUrl(),
AuthHeaderName: req.GetAuthHeaderName(),
AuthHeaderValue: req.GetAuthHeaderValue(),
SkipTlsVerify: req.GetSkipTlsVerify(),
OllamaFallback: req.GetOllamaFallback(),
}
for _, proxyID := range proxyIDs {
connVal, ok := s.connectedProxies.Load(proxyID)
if !ok {
continue
}
conn := connVal.(*proxyConnection)
// Cluster membership can briefly retain a proxy ID from a superseded
// session. The live connection is authoritative: never send provider
// credentials unless it is connected to the exact requested cluster.
if conn.address != clusterAddr {
continue
}
if !modelDiscoveryCapable(conn, accountID) {
continue
}
pending := &pendingModelDiscovery{
proxyID: conn.proxyID,
sessionID: conn.sessionID,
done: make(chan modelDiscoveryCompletion, 1),
}
s.registerModelDiscovery(request.GetRequestId(), pending)
queued := false
select {
case conn.modelDiscoveryChan <- request:
queued = true
case <-ctx.Done():
s.removeModelDiscovery(request.GetRequestId(), pending)
return nil, modelDiscoveryContextError(ctx.Err())
default:
s.removeModelDiscovery(request.GetRequestId(), pending)
}
if !queued {
continue
}
defer s.removeModelDiscovery(request.GetRequestId(), pending)
select {
case completion := <-pending.done:
if completion.err != nil {
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery proxy disconnected")
}
return completion.result, nil
case <-ctx.Done():
return nil, modelDiscoveryContextError(ctx.Err())
case <-conn.ctx.Done():
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery proxy disconnected")
}
}
return nil, nbstatus.Errorf(nbstatus.PreconditionFailed, "no connected proxy in the configured cluster supports model discovery")
}
func modelDiscoveryCapable(conn *proxyConnection, accountID string) bool {
if conn == nil || conn.syncStream == nil || conn.modelDiscoveryChan == nil || conn.ready == nil {
return false
}
if conn.ctx == nil || conn.ctx.Err() != nil {
return false
}
if conn.accountID != nil && *conn.accountID != accountID {
return false
}
if conn.capabilities == nil || !conn.capabilities.GetSupportsModelDiscovery() {
return false
}
select {
case <-conn.ready:
return true
default:
return false
}
}
func modelDiscoveryContextError(err error) error {
if errors.Is(err, context.DeadlineExceeded) {
return nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery timed out")
}
return nbstatus.Errorf(nbstatus.PreconditionFailed, "model discovery was canceled")
}
func (s *ProxyServiceServer) registerModelDiscovery(requestID string, pending *pendingModelDiscovery) {
s.modelDiscoveryMu.Lock()
defer s.modelDiscoveryMu.Unlock()
if s.modelDiscoveryPending == nil {
s.modelDiscoveryPending = make(map[string]*pendingModelDiscovery)
}
s.modelDiscoveryPending[requestID] = pending
}
func (s *ProxyServiceServer) removeModelDiscovery(requestID string, pending *pendingModelDiscovery) {
s.modelDiscoveryMu.Lock()
defer s.modelDiscoveryMu.Unlock()
if s.modelDiscoveryPending[requestID] == pending {
delete(s.modelDiscoveryPending, requestID)
}
}
func (s *ProxyServiceServer) isModelDiscoveryPendingFor(conn *proxyConnection, requestID string) bool {
if conn == nil || requestID == "" {
return false
}
s.modelDiscoveryMu.Lock()
defer s.modelDiscoveryMu.Unlock()
pending := s.modelDiscoveryPending[requestID]
return pending != nil && pending.proxyID == conn.proxyID && pending.sessionID == conn.sessionID
}
func (s *ProxyServiceServer) completeModelDiscovery(conn *proxyConnection, result *proto.ModelDiscoveryResult) {
if conn == nil || result == nil || result.GetRequestId() == "" {
return
}
s.modelDiscoveryMu.Lock()
pending := s.modelDiscoveryPending[result.GetRequestId()]
s.modelDiscoveryMu.Unlock()
if pending == nil || pending.proxyID != conn.proxyID || pending.sessionID != conn.sessionID {
return
}
select {
case pending.done <- modelDiscoveryCompletion{result: result}:
default:
}
}
func (s *ProxyServiceServer) failModelDiscoveries(conn *proxyConnection, err error) {
if conn == nil {
return
}
s.modelDiscoveryMu.Lock()
pending := make([]*pendingModelDiscovery, 0)
for _, item := range s.modelDiscoveryPending {
if item.proxyID == conn.proxyID && item.sessionID == conn.sessionID {
pending = append(pending, item)
}
}
s.modelDiscoveryMu.Unlock()
for _, item := range pending {
select {
case item.done <- modelDiscoveryCompletion{err: err}:
default:
}
}
}
// SendServiceUpdateToCluster sends a service update to all proxy servers in a specific cluster.
// If clusterAddr is empty, broadcasts to all connected proxy servers (backward compatibility).
// For create/update operations a unique one-time auth token is generated per
@@ -1304,12 +1049,6 @@ func (s *ProxyServiceServer) SendServiceUpdateToCluster(ctx context.Context, upd
continue
}
conn := connVal.(*proxyConnection)
// Membership can retain this proxy ID from a superseded session in a
// different cluster. The live connection address is authoritative,
// especially because Agent Network mappings can contain credentials.
if conn.address != clusterAddr {
continue
}
if conn.accountID != nil && update.AccountId != "" && *conn.accountID != update.AccountId {
continue
}

View File

@@ -42,10 +42,6 @@ func newTestProxyController() *testProxyController {
func (c *testProxyController) SendServiceUpdateToCluster(_ context.Context, _ string, _ *proto.ProxyMapping, _ string) {
}
func (c *testProxyController) DiscoverModels(_ context.Context, _, _ string, _ *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
return nil, proxy.ErrModelDiscoveryUnavailable
}
func (c *testProxyController) GetOIDCValidationConfig() proxy.OIDCValidationConfig {
return proxy.OIDCValidationConfig{}
}
@@ -222,65 +218,6 @@ func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
assert.Empty(t, msg2.AuthToken)
}
func TestSendServiceUpdateToCluster_RejectsStaleClusterMembership(t *testing.T) {
ctx := context.Background()
s := &ProxyServiceServer{
tokenStore: NewOneTimeTokenStore(ctx, testCacheStore(t)),
}
controller := newTestProxyController()
s.SetProxyController(controller)
ch := registerFakeProxy(s, "proxy-a", "new-cluster.example.com")
require.NoError(t, controller.RegisterProxyToCluster(ctx, "old-cluster.example.com", "proxy-a"))
s.SendServiceUpdateToCluster(ctx, &proto.ProxyMapping{
Type: proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED,
Id: "agent-network-provider-1",
AccountId: "account-1",
Domain: "agent.example.com",
Path: []*proto.PathMapping{
{Path: "/", Target: "http://ollama.internal:11434/"},
},
}, "old-cluster.example.com")
assert.True(t, drainEmpty(ch), "stale membership must not route a mapping to a different live cluster")
}
func TestDisconnectProxy_RemovesSupersededMembershipFromOldCluster(t *testing.T) {
controller := newTestProxyController()
server := &ProxyServiceServer{proxyController: controller}
oldCtx, cancelOld := context.WithCancel(context.Background())
newCtx, cancelNew := context.WithCancel(context.Background())
t.Cleanup(cancelOld)
t.Cleanup(cancelNew)
oldConnection := &proxyConnection{
proxyID: "proxy-a",
sessionID: "old-session",
address: "old-cluster.example.com",
ctx: oldCtx,
cancel: cancelOld,
}
newConnection := &proxyConnection{
proxyID: "proxy-a",
sessionID: "new-session",
address: "new-cluster.example.com",
ctx: newCtx,
cancel: cancelNew,
}
require.NoError(t, controller.RegisterProxyToCluster(context.Background(), oldConnection.address, oldConnection.proxyID))
require.NoError(t, controller.RegisterProxyToCluster(context.Background(), newConnection.address, newConnection.proxyID))
server.connectedProxies.Store(newConnection.proxyID, newConnection)
server.disconnectProxy(oldConnection)
assert.Empty(t, controller.GetProxiesForCluster(oldConnection.address))
assert.Equal(t, []string{"proxy-a"}, controller.GetProxiesForCluster(newConnection.address))
current, ok := server.connectedProxies.Load(newConnection.proxyID)
require.True(t, ok)
assert.Same(t, newConnection, current)
}
func TestSendServiceUpdate_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))

View File

@@ -6,7 +6,6 @@ import (
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
"github.com/netbirdio/netbird/management/server/permissions"
"github.com/netbirdio/netbird/management/server/permissions/modules"
"github.com/netbirdio/netbird/management/server/permissions/operations"
@@ -31,10 +30,6 @@ type managerImpl struct {
accountManager account.Manager
}
func eventMetaResource(group *types.Group, resource *resourceTypes.NetworkResource) map[string]any {
return map[string]any{"name": group.Name, "id": group.ID, "resource_name": resource.Name, "resource_id": resource.ID, "resource_type": resource.Type}
}
type mockManager struct {
}
@@ -114,7 +109,7 @@ func (m *managerImpl) AddResourceToGroupInTransaction(ctx context.Context, trans
}
event := func() {
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, eventMetaResource(group, networkResource))
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceAddedToGroup, group.EventMetaResource(networkResource))
}
return event, nil
@@ -138,7 +133,7 @@ func (m *managerImpl) RemoveResourceFromGroupInTransaction(ctx context.Context,
}
event := func() {
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, eventMetaResource(group, networkResource))
m.accountManager.StoreEvent(ctx, userID, groupID, accountID, activity.ResourceRemovedFromGroup, group.EventMetaResource(networkResource))
}
return event, nil

View File

@@ -446,7 +446,7 @@ func (h *Handler) GetAccessiblePeers(w http.ResponseWriter, r *http.Request) {
netMap := account.GetPeerNetworkMapFromComponents(ctx, peerID, dns.CustomZone{}, nil, validPeers, account.GetResourcePoliciesMap(), account.GetResourceRoutersMap(), nil, account.GetActiveGroupUsers())
util.WriteJSONObject(ctx, w, toAccessiblePeers(netMap, account.Peers, dnsDomain))
util.WriteJSONObject(ctx, w, toAccessiblePeers(netMap, dnsDomain))
}
func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request) {
@@ -534,20 +534,15 @@ func (h *Handler) CreateTemporaryAccess(w http.ResponseWriter, r *http.Request)
util.WriteJSONObject(r.Context(), w, resp)
}
// toAccessiblePeers rehydrates the calculated map's component peers into the
// account's full peer objects, which carry the location/status/meta fields
// the API response needs.
func toAccessiblePeers(netMap *types.NetworkMap, accountPeers map[string]*nbpeer.Peer, dnsDomain string) []api.AccessiblePeer {
func toAccessiblePeers(netMap *types.NetworkMap, dnsDomain string) []api.AccessiblePeer {
accessiblePeers := make([]api.AccessiblePeer, 0, len(netMap.Peers)+len(netMap.OfflinePeers))
add := func(peers []*types.ComponentPeer) {
for _, p := range peers {
if peer := accountPeers[p.ID]; peer != nil {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(peer, dnsDomain))
}
}
for _, p := range netMap.Peers {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
for _, p := range netMap.OfflinePeers {
accessiblePeers = append(accessiblePeers, peerToAccessiblePeer(p, dnsDomain))
}
add(netMap.Peers)
add(netMap.OfflinePeers)
return accessiblePeers
}

View File

@@ -683,81 +683,3 @@ func BackfillPublicIDs[T any](ctx context.Context, db *gorm.DB) error {
log.WithContext(ctx).Infof("Backfill of empty public_id in table %s completed", tableName)
return nil
}
// FoldCostAggregatesIntoBuckets migrates a per-request cost table from the old
// "stored aggregate" shape (cost_usd + cache_cost_usd columns) to the per-bucket
// breakdown, where the total and cache portion are derived on read instead.
//
// The fold preserves both aggregates exactly for historical rows: the cache
// total moves into cached_input_cost_usd and the remainder into
// input_cost_usd, so a row's derived total and cache cost still match what it
// reported before the upgrade. The finer split is genuinely unknown for those
// rows — the old schema never recorded a read/write or input/output division —
// so it is lumped rather than guessed; only rows written after the upgrade
// carry a true four-way split.
//
// Dropping the columns before folding would zero every historical row's cost,
// so the update runs first and the drop only happens once it succeeds. A table
// with no cost_usd column has already been migrated (or was created fresh) and
// is skipped.
func FoldCostAggregatesIntoBuckets[T any](ctx context.Context, db *gorm.DB) error {
var model T
if !db.Migrator().HasTable(&model) {
log.WithContext(ctx).Debugf("table for %T does not exist, no cost-bucket migration needed", model)
return nil
}
if !db.Migrator().HasColumn(&model, "cost_usd") {
log.WithContext(ctx).Debugf("table for %T has no cost_usd column, cost buckets already migrated", model)
return nil
}
stmt := &gorm.Statement{DB: db}
if err := stmt.Parse(&model); err != nil {
return fmt.Errorf("parse model schema: %w", err)
}
tableName := stmt.Schema.Table
// COALESCE guards rows whose new columns were added as NULL by an earlier
// AutoMigrate run that predates the NOT NULL default.
hasCacheColumn := db.Migrator().HasColumn(&model, "cache_cost_usd")
cacheExpr := "0"
if hasCacheColumn {
cacheExpr = "COALESCE(cache_cost_usd, 0)"
}
if err := db.Transaction(func(tx *gorm.DB) error {
// Only touch rows that carry a legacy total and no breakdown yet, so
// the migration is idempotent and never overwrites a true split.
update := fmt.Sprintf(`UPDATE %s
SET input_cost_usd = COALESCE(cost_usd, 0) - %s,
cached_input_cost_usd = %s,
cache_creation_cost_usd = 0,
output_cost_usd = 0
WHERE COALESCE(cost_usd, 0) <> 0
AND COALESCE(input_cost_usd, 0) = 0
AND COALESCE(cached_input_cost_usd, 0) = 0
AND COALESCE(cache_creation_cost_usd, 0) = 0
AND COALESCE(output_cost_usd, 0) = 0`, tableName, cacheExpr, cacheExpr)
res := tx.Exec(update)
if res.Error != nil {
return fmt.Errorf("fold legacy cost aggregates in %s: %w", tableName, res.Error)
}
log.WithContext(ctx).Infof("folded legacy cost aggregates into per-bucket columns for %d rows in table %s", res.RowsAffected, tableName)
if err := tx.Migrator().DropColumn(&model, "cost_usd"); err != nil {
return fmt.Errorf("drop cost_usd from %s: %w", tableName, err)
}
if hasCacheColumn {
if err := tx.Migrator().DropColumn(&model, "cache_cost_usd"); err != nil {
return fmt.Errorf("drop cache_cost_usd from %s: %w", tableName, err)
}
}
return nil
}); err != nil {
return err
}
log.WithContext(ctx).Infof("migration of stored cost aggregates to per-bucket columns in table %s completed", tableName)
return nil
}

View File

@@ -16,7 +16,6 @@ import (
"gorm.io/driver/sqlite"
"gorm.io/gorm"
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
"github.com/netbirdio/netbird/management/server/migration"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/testutil"
@@ -640,99 +639,3 @@ func TestCleanupOrphanedResources_SkipsWhenForeignKeyExists(t *testing.T) {
db.Model(&testChildWithFK{}).Count(&count)
assert.Equal(t, int64(2), count, "Both rows should survive — migration must skip when FK constraint exists")
}
// legacyCostRow is the pre-breakdown shape of the usage table: cost was stored
// as a total plus a cache portion, with no per-bucket columns. Used to build a
// realistic pre-upgrade table for the fold migration to run against.
type legacyCostRow struct {
ID string `gorm:"primaryKey"`
AccountID string
Model string
CostUSD float64
CacheCostUSD float64
}
func (legacyCostRow) TableName() string { return "agent_network_request_usage" }
// TestFoldCostAggregatesIntoBuckets_PreservesHistoricalCost covers the upgrade
// path: a table written under the old schema must come out with its per-row
// total and cache cost unchanged, because dropping cost_usd without folding it
// forward would silently zero every historical row's spend.
func TestFoldCostAggregatesIntoBuckets_PreservesHistoricalCost(t *testing.T) {
ctx := context.Background()
db := setupDatabase(t)
// setupDatabase hands back a process-shared database, so start from a clean
// table rather than inheriting rows from another test.
require.NoError(t, db.Migrator().DropTable(&agentNetworkTypes.AgentNetworkUsage{}))
require.NoError(t, db.AutoMigrate(&legacyCostRow{}), "legacy table must be created")
require.NoError(t, db.Create(&legacyCostRow{
ID: "u1", AccountID: "acct-1", Model: "claude-sonnet-4-6", CostUSD: 0.0123, CacheCostUSD: 0.0029,
}).Error)
require.NoError(t, db.Create(&legacyCostRow{
ID: "u2", AccountID: "acct-1", Model: "gpt-4o", CostUSD: 0.5, CacheCostUSD: 0,
}).Error)
// A zero-cost row (denied / unpriced request) must stay zero, not be touched.
require.NoError(t, db.Create(&legacyCostRow{ID: "u3", AccountID: "acct-1", Model: "gw/unpriced"}).Error)
// AutoMigrate adds the per-bucket columns alongside the legacy ones, exactly
// as a real upgrade does before the post-auto migrations run.
require.NoError(t, db.AutoMigrate(&agentNetworkTypes.AgentNetworkUsage{}), "new columns must be added")
require.NoError(t, migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkUsage](ctx, db))
assert.False(t, db.Migrator().HasColumn(&agentNetworkTypes.AgentNetworkUsage{}, "cost_usd"),
"legacy cost_usd column must be dropped once folded")
assert.False(t, db.Migrator().HasColumn(&agentNetworkTypes.AgentNetworkUsage{}, "cache_cost_usd"),
"legacy cache_cost_usd column must be dropped once folded")
var rows []*agentNetworkTypes.AgentNetworkUsage
require.NoError(t, db.Order("id").Find(&rows).Error)
require.Len(t, rows, 3)
// u1: total and cache portion both preserved; the read/write and
// input/output splits are unknowable for a legacy row, so the cache total
// lands on cached_input and the remainder on input.
assert.InDelta(t, 0.0123, rows[0].TotalCostUSD(), 1e-9, "historical total must survive the fold")
assert.InDelta(t, 0.0029, rows[0].CacheCostUSD(), 1e-9, "historical cache cost must survive the fold")
assert.InDelta(t, 0.0094, rows[0].InputCostUSD, 1e-9, "non-cache remainder lands on input")
assert.InDelta(t, 0.0029, rows[0].CachedInputCostUSD, 1e-9, "legacy cache total lands on cached input")
assert.Zero(t, rows[0].CacheCreationCostUSD, "legacy rows carry no read/write split to recover")
assert.Zero(t, rows[0].OutputCostUSD, "legacy rows carry no input/output split to recover")
// u2: no cache spend — the whole total is the non-cache remainder.
assert.InDelta(t, 0.5, rows[1].TotalCostUSD(), 1e-9, "cache-free historical total must survive")
assert.Zero(t, rows[1].CacheCostUSD(), "a cache-free row must stay cache-free")
// u3: zero stays zero rather than being rewritten.
assert.Zero(t, rows[2].TotalCostUSD(), "an unpriced row must remain unpriced")
}
// TestFoldCostAggregatesIntoBuckets_SkipsAlreadyMigrated proves the migration is
// safe to re-run: with no legacy column present it is a no-op that leaves a
// true four-way split untouched.
func TestFoldCostAggregatesIntoBuckets_SkipsAlreadyMigrated(t *testing.T) {
ctx := context.Background()
db := setupDatabase(t)
require.NoError(t, db.Migrator().DropTable(&agentNetworkTypes.AgentNetworkUsage{}))
require.NoError(t, db.AutoMigrate(&agentNetworkTypes.AgentNetworkUsage{}))
// Timestamp must be set explicitly: a zero time.Time serialises as
// '0000-00-00 00:00:00', which MySQL rejects under strict mode.
require.NoError(t, db.Create(&agentNetworkTypes.AgentNetworkUsage{
ID: "u1", AccountID: "acct-1", Model: "claude-sonnet-4-6",
Timestamp: time.Date(2026, 5, 5, 9, 0, 0, 0, time.UTC),
InputCostUSD: 0.001, CachedInputCostUSD: 0.002, CacheCreationCostUSD: 0.003, OutputCostUSD: 0.004,
}).Error)
require.NoError(t, migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkUsage](ctx, db),
"running against an already-migrated table must be a no-op, not an error")
var row agentNetworkTypes.AgentNetworkUsage
require.NoError(t, db.First(&row, "id = ?", "u1").Error)
assert.InDelta(t, 0.001, row.InputCostUSD, 1e-9, "a true split must not be rewritten")
assert.InDelta(t, 0.002, row.CachedInputCostUSD, 1e-9)
assert.InDelta(t, 0.003, row.CacheCreationCostUSD, 1e-9)
assert.InDelta(t, 0.004, row.OutputCostUSD, 1e-9)
assert.InDelta(t, 0.01, row.TotalCostUSD(), 1e-9, "derived total sums the four buckets")
}

View File

@@ -14,7 +14,6 @@ import (
nbDomain "github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
type NetworkResourceType string
@@ -65,27 +64,6 @@ func NewNetworkResource(accountID, networkID, name, description, address string,
}, nil
}
// ToComponent converts the resource to its self-contained components
// representation. Returns nil for a nil resource.
func (n *NetworkResource) ToComponent() *sharedTypes.ComponentResource {
if n == nil {
return nil
}
return &sharedTypes.ComponentResource{
ID: n.ID,
PublicID: n.PublicID,
NetworkID: n.NetworkID,
AccountID: n.AccountID,
Name: n.Name,
Description: n.Description,
Type: sharedTypes.ComponentResourceType(n.Type),
Address: n.Address,
Domain: n.Domain,
Prefix: n.Prefix,
Enabled: n.Enabled,
}
}
func (n *NetworkResource) ToAPIResponse(groups []api.GroupMinimum) *api.NetworkResource {
addr := n.Prefix.String()
if n.Type == Domain {

View File

@@ -7,7 +7,6 @@ import (
"github.com/netbirdio/netbird/management/server/networks/types"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
type NetworkRouter struct {
@@ -22,36 +21,6 @@ type NetworkRouter struct {
Enabled bool
}
// ToComponent converts the router to its self-contained components
// representation. Returns nil for a nil router.
func (n *NetworkRouter) ToComponent() *sharedTypes.ComponentRouter {
if n == nil {
return nil
}
return &sharedTypes.ComponentRouter{
NetworkID: n.NetworkID,
PublicID: n.PublicID,
Peer: n.Peer,
PeerGroups: n.PeerGroups,
Masquerade: n.Masquerade,
Metric: n.Metric,
Enabled: n.Enabled,
}
}
// ToComponentMap converts a peer-keyed router map to its components
// representation.
func ToComponentMap(routers map[string]*NetworkRouter) map[string]*sharedTypes.ComponentRouter {
if routers == nil {
return nil
}
out := make(map[string]*sharedTypes.ComponentRouter, len(routers))
for id, r := range routers {
out[id] = r.ToComponent()
}
return out
}
func NewNetworkRouter(accountID string, networkID string, peer string, peerGroups []string, masquerade bool, metric int, enabled bool) (*NetworkRouter, error) {
r := &NetworkRouter{
ID: xid.New().String(),

View File

@@ -405,7 +405,7 @@ func (am *DefaultAccountManager) CreatePeerJob(ctx context.Context, accountID, p
return status.NewPeerNotPartOfAccountError()
}
meetMinVer, err := version.MeetsMinVersion(remoteJobsMinVer, p.Meta.WtVersion)
meetMinVer, err := posture.MeetsMinVersion(remoteJobsMinVer, p.Meta.WtVersion)
if !version.IsDevelopmentVersion(p.Meta.WtVersion) && (!meetMinVer || err != nil) {
return status.Errorf(status.PreconditionFailed, "peer version %s does not meet the minimum required version %s for remote jobs", p.Meta.WtVersion, remoteJobsMinVer)
}
@@ -1588,7 +1588,7 @@ func affectedPeerIDsFromNetworkMap(nmap *types.NetworkMap, selfPeerID string) []
}
seen := make(map[string]struct{}, len(nmap.Peers)+len(nmap.OfflinePeers))
ids := make([]string, 0, len(nmap.Peers)+len(nmap.OfflinePeers))
add := func(peers []*types.ComponentPeer) {
add := func(peers []*nbpeer.Peer) {
for _, p := range peers {
if p == nil || p.ID == "" || p.ID == selfPeerID {
continue

View File

@@ -13,7 +13,6 @@ import (
"github.com/netbirdio/netbird/management/server/util"
"github.com/netbirdio/netbird/shared/management/http/api"
sharedTypes "github.com/netbirdio/netbird/shared/management/types"
)
// Peer capability constants mirror the proto enum values.
@@ -206,35 +205,6 @@ func (p *Peer) AddedWithSSOLogin() bool {
return p.UserID != ""
}
// ToComponent converts the peer to its self-contained components
// representation, carrying exactly the subset of peer data that crosses the
// components wire format. Returns nil for a nil peer so callers can convert
// possibly-missing peers without guarding.
func (p *Peer) ToComponent() *sharedTypes.ComponentPeer {
if p == nil {
return nil
}
cp := &sharedTypes.ComponentPeer{
ID: p.ID,
Key: p.Key,
IP: p.IP,
IPv6: p.IPv6,
DNSLabel: p.DNSLabel,
SSHKey: p.SSHKey,
SSHEnabled: p.SSHEnabled,
ServerSSHAllowed: p.Meta.Flags.ServerSSHAllowed,
AgentVersion: p.Meta.WtVersion,
SupportsSourcePrefixes: p.SupportsSourcePrefixes(),
SupportsIPv6: p.SupportsIPv6(),
LoginExpirationEnabled: p.LoginExpirationEnabled,
AddedWithSSOLogin: p.AddedWithSSOLogin(),
}
if p.LastLogin != nil {
cp.LastLogin = *p.LastLogin
}
return cp
}
// HasCapability reports whether the peer has the given capability.
func (p *Peer) HasCapability(capability int32) bool {
return slices.Contains(p.Meta.Capabilities, capability)

View File

@@ -1092,14 +1092,14 @@ func TestToSyncResponse(t *testing.T) {
}
networkMap := &types.NetworkMap{
Network: &types.Network{Net: *ipnet, Serial: 1000},
Peers: []*types.ComponentPeer{{
Peers: []*nbpeer.Peer{{
IP: netip.MustParseAddr("192.168.1.2"),
IPv6: netip.MustParseAddr("fd00::2"),
Key: "peer2-key",
DNSLabel: "peer2",
SSHEnabled: true,
SSHKey: "peer2-ssh-key"}},
OfflinePeers: []*types.ComponentPeer{{
OfflinePeers: []*nbpeer.Peer{{
IP: netip.MustParseAddr("192.168.1.3"),
IPv6: netip.MustParseAddr("fd00::3"),
Key: "peer3-key",

View File

@@ -3,9 +3,11 @@ package posture
import (
"context"
"fmt"
"strings"
"github.com/hashicorp/go-version"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
nbversion "github.com/netbirdio/netbird/version"
)
type NBVersionCheck struct {
@@ -14,8 +16,14 @@ type NBVersionCheck struct {
var _ Check = (*NBVersionCheck)(nil)
// sanitizeVersion removes anything after the pre-release tag (e.g., "-dev", "-alpha", etc.)
func sanitizeVersion(version string) string {
parts := strings.Split(version, "-")
return parts[0]
}
func (n *NBVersionCheck) Check(ctx context.Context, peer nbpeer.Peer) (bool, error) {
meetsMin, err := nbversion.MeetsMinVersion(n.MinVersion, peer.Meta.WtVersion)
meetsMin, err := MeetsMinVersion(n.MinVersion, peer.Meta.WtVersion)
if err != nil {
return false, err
}
@@ -40,3 +48,21 @@ func (n *NBVersionCheck) Validate() error {
}
return nil
}
// MeetsMinVersion checks if the peer's version meets or exceeds the minimum required version
func MeetsMinVersion(minVer, peerVer string) (bool, error) {
peerVer = sanitizeVersion(peerVer)
minVer = sanitizeVersion(minVer)
peerNBVer, err := version.NewVersion(peerVer)
if err != nil {
return false, err
}
constraints, err := version.NewConstraint(">= " + minVer)
if err != nil {
return false, err
}
return constraints.Check(peerNBVer), nil
}

View File

@@ -139,3 +139,68 @@ func TestNBVersionCheck_Validate(t *testing.T) {
})
}
}
func TestMeetsMinVersion(t *testing.T) {
tests := []struct {
name string
minVer string
peerVer string
want bool
wantErr bool
}{
{
name: "Peer version greater than min version",
minVer: "0.26.0",
peerVer: "0.60.1",
want: true,
wantErr: false,
},
{
name: "Peer version equals min version",
minVer: "1.0.0",
peerVer: "1.0.0",
want: true,
wantErr: false,
},
{
name: "Peer version less than min version",
minVer: "1.0.0",
peerVer: "0.9.9",
want: false,
wantErr: false,
},
{
name: "Peer version with pre-release tag greater than min version",
minVer: "1.0.0",
peerVer: "1.0.1-alpha",
want: true,
wantErr: false,
},
{
name: "Invalid peer version format",
minVer: "1.0.0",
peerVer: "dev",
want: false,
wantErr: true,
},
{
name: "Invalid min version format",
minVer: "invalid.version",
peerVer: "1.0.0",
want: false,
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := MeetsMinVersion(tt.minVer, tt.peerVer)
if tt.wantErr {
assert.Error(t, err)
} else {
assert.NoError(t, err)
}
assert.Equal(t, tt.want, got)
})
}
}

View File

@@ -71,7 +71,7 @@ func (s *SqlStore) GetAgentNetworkMetrics(ctx context.Context) (AgentNetworkMetr
usageRow := db.Model(&agentNetworkTypes.AgentNetworkUsage{}).
Select("COALESCE(SUM(input_tokens), 0) AS input_tokens, " +
"COALESCE(SUM(output_tokens), 0) AS output_tokens, " +
"COALESCE(SUM" + agentNetworkTypes.CostUSDSQLExpr + ", 0) AS cost_usd").Row()
"COALESCE(SUM(cost_usd), 0) AS cost_usd").Row()
if err := usageRow.Scan(&m.InputTokens, &m.OutputTokens, &m.CostUSD); err != nil {
return AgentNetworkMetrics{}, fmt.Errorf("scan agent network usage metrics: %w", err)
}

View File

@@ -37,7 +37,7 @@ func TestAgentNetworkUsage_RealStore_RoundTrip(t *testing.T) {
InputTokens: 1200,
OutputTokens: 640,
TotalTokens: 1840,
InputCostUSD: 0.0231,
CostUSD: 0.0231,
}
usageGroups := []agentNetworkTypes.AgentNetworkUsageGroup{
{UsageID: usage.ID, GroupID: "grp-eng", AccountID: accountID},
@@ -71,7 +71,7 @@ func TestAgentNetworkUsage_RealStore_RoundTrip(t *testing.T) {
InputTokens: 1200,
OutputTokens: 640,
TotalTokens: 1840,
InputCostUSD: 0.0231,
CostUSD: 0.0231,
}
entryGroups := []agentNetworkTypes.AgentNetworkAccessLogGroup{
{LogID: entry.ID, GroupID: "grp-eng", AccountID: accountID},
@@ -127,7 +127,7 @@ func TestAgentNetworkUsageOverview_DailyAggregation(t *testing.T) {
mk := func(id string, ts time.Time, model string, in, out int64, cost float64) *agentNetworkTypes.AgentNetworkUsage {
return &agentNetworkTypes.AgentNetworkUsage{
ID: id, AccountID: accountID, Timestamp: ts, Model: model,
InputTokens: in, OutputTokens: out, TotalTokens: in + out, InputCostUSD: cost,
InputTokens: in, OutputTokens: out, TotalTokens: in + out, CostUSD: cost,
}
}
require.NoError(t, s.CreateAgentNetworkUsage(ctx, mk("u1", day1, "gpt-4o", 100, 50, 0.10), nil))
@@ -143,7 +143,7 @@ func TestAgentNetworkUsageOverview_DailyAggregation(t *testing.T) {
assert.Equal(t, "2026-05-05", buckets[0].PeriodStart, "oldest-first ordering")
assert.Equal(t, int64(300), buckets[0].InputTokens, "same-day input tokens summed")
assert.Equal(t, int64(130), buckets[0].OutputTokens)
assert.InDelta(t, 0.30, buckets[0].TotalCostUSD(), 1e-9, "same-day cost summed")
assert.InDelta(t, 0.30, buckets[0].CostUSD, 1e-9, "same-day cost summed")
assert.Equal(t, "2026-05-06", buckets[1].PeriodStart)
assert.Equal(t, int64(15), buckets[1].TotalTokens)
@@ -174,7 +174,7 @@ func TestAgentNetworkAccessLogSessions_RealStore(t *testing.T) {
ID: id, AccountID: accountID, ServiceID: "svc", Timestamp: ts,
UserID: user, StatusCode: 200, Provider: provider, Model: model,
SessionID: session, Decision: decision,
InputTokens: 100, OutputTokens: 50, TotalTokens: 150, InputCostUSD: cost,
InputTokens: 100, OutputTokens: 50, TotalTokens: 150, CostUSD: cost,
}
}
@@ -207,7 +207,7 @@ func TestAgentNetworkAccessLogSessions_RealStore(t *testing.T) {
s1 := sessions[2]
assert.Equal(t, 2, s1.RequestCount, "s1 has two requests")
assert.Equal(t, int64(300), s1.TotalTokens, "tokens summed across the session")
assert.InDelta(t, 0.30, s1.TotalCostUSD(), 1e-9, "cost summed across the session")
assert.InDelta(t, 0.30, s1.CostUSD, 1e-9, "cost summed across the session")
assert.Equal(t, "alice", s1.UserID)
assert.Equal(t, "allow", s1.Decision)
// SQLite hands times back in time.Local; normalise to UTC so the instant is

View File

@@ -650,14 +650,6 @@ func getMigrationsPostAuto(ctx context.Context) []migrationFunc {
func(db *gorm.DB) error {
return migration.DropIndex[proxy.Proxy](ctx, db, "idx_proxy_account_id_unique")
},
// Post-auto so the per-bucket cost columns already exist when the legacy
// aggregates are folded into them and dropped.
func(db *gorm.DB) error {
return migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkAccessLog](ctx, db)
},
func(db *gorm.DB) error {
return migration.FoldCostAggregatesIntoBuckets[agentNetworkTypes.AgentNetworkUsage](ctx, db)
},
}
}

View File

@@ -1082,7 +1082,6 @@ func (a *Account) connResourcesGenerator(ctx context.Context, targetPeer *nbpeer
peersExists := make(map[string]struct{})
rules := make([]*FirewallRule, 0)
peers := make([]*nbpeer.Peer, 0)
targetComponent := targetPeer.ToComponent()
return func(rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int) {
for _, peer := range groupPeers {
@@ -1118,10 +1117,10 @@ func (a *Account) connResourcesGenerator(ctx context.Context, targetPeer *nbpeer
if len(rule.Ports) == 0 && len(rule.PortRanges) == 0 {
rules = append(rules, &fr)
} else {
rules = append(rules, ExpandPortsAndRanges(fr, rule, targetComponent)...)
rules = append(rules, ExpandPortsAndRanges(fr, rule, targetPeer)...)
}
rules = AppendIPv6FirewallRule(rules, rulesExists, peer.ToComponent(), targetComponent, rule, FirewallRuleContext{
rules = AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, FirewallRuleContext{
Direction: direction,
DirStr: strconv.Itoa(direction),
ProtocolStr: string(protocol),
@@ -1281,7 +1280,7 @@ func (a *Account) getRouteFirewallRules(ctx context.Context, peerID string, poli
return fwRules
}
func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}, validatedPeersMap map[string]struct{}) []*ComponentPeer {
func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}, validatedPeersMap map[string]struct{}) []*nbpeer.Peer {
distPeersWithPolicy := make(map[string]struct{})
for _, id := range rule.Sources {
group := a.Groups[id]
@@ -1308,13 +1307,13 @@ func (a *Account) getRulePeers(rule *PolicyRule, postureChecks []string, peerID
}
}
distributionGroupPeers := make([]*ComponentPeer, 0, len(distPeersWithPolicy))
distributionGroupPeers := make([]*nbpeer.Peer, 0, len(distPeersWithPolicy))
for pID := range distPeersWithPolicy {
peer := a.Peers[pID]
if peer == nil {
continue
}
distributionGroupPeers = append(distributionGroupPeers, peer.ToComponent())
distributionGroupPeers = append(distributionGroupPeers, peer)
}
return distributionGroupPeers
}

View File

@@ -9,7 +9,9 @@ import (
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/management/internals/modules/zones"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/route"
)
@@ -111,7 +113,7 @@ func (a *Account) GetPeerNetworkMapComponents(
PeerID: peerID,
Network: a.Network.Copy(),
// must include the target peer as it's required on the client
Peers: map[string]*ComponentPeer{peerID: peer.ToComponent()},
Peers: map[string]*nbpeer.Peer{peerID: peer},
})
}
@@ -124,7 +126,7 @@ func (a *Account) GetPeerNetworkMapComponents(
PeerID: peerID,
Network: a.Network.Copy(),
// must include the target peer as it's required on the client
Peers: map[string]*ComponentPeer{peerID: peer.ToComponent()},
Peers: map[string]*nbpeer.Peer{peerID: peer},
})
}
@@ -134,10 +136,10 @@ func (a *Account) GetPeerNetworkMapComponents(
NameServerGroups: make([]*nbdns.NameServerGroup, 0),
CustomZoneDomain: peersCustomZone.Domain,
ResourcePoliciesMap: make(map[string][]*Policy),
RoutersMap: make(map[string]map[string]*ComponentRouter),
NetworkResources: make([]*ComponentResource, 0),
RoutersMap: make(map[string]map[string]*routerTypes.NetworkRouter),
NetworkResources: make([]*resourceTypes.NetworkResource, 0),
PostureFailedPeers: make(map[string]map[string]struct{}, len(a.PostureChecks)),
RouterPeers: make(map[string]*ComponentPeer),
RouterPeers: make(map[string]*nbpeer.Peer),
NetworkXIDToPublicID: make(map[string]string, len(a.Networks)),
PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)),
}
@@ -172,7 +174,7 @@ func (a *Account) GetPeerNetworkMapComponents(
}
components.Peers = relevantPeers
components.Groups = GroupsToComponent(relevantGroups)
components.Groups = relevantGroups
components.Policies = relevantPolicies
components.Routes = relevantRoutes
components.AllDNSRecords = filterDNSRecordsByPeers(peersCustomZone.Records, relevantPeers, peer.SupportsIPv6() && peer.IPv6.IsValid())
@@ -221,7 +223,7 @@ func (a *Account) GetPeerNetworkMapComponents(
}
for _, pID := range a.getPostureValidPeersSaveFailed(peers, policy.SourcePostureChecks, validatedPeersMap, &components.PostureFailedPeers) {
if _, exists := components.Peers[pID]; !exists {
components.Peers[pID] = a.GetPeer(pID).ToComponent()
components.Peers[pID] = a.GetPeer(pID)
}
}
} else {
@@ -254,14 +256,14 @@ func (a *Account) GetPeerNetworkMapComponents(
for _, srcGroupID := range rule.Sources {
if g := a.Groups[srcGroupID]; g != nil {
if _, exists := components.Groups[srcGroupID]; !exists {
components.Groups[srcGroupID] = g.ToComponent()
components.Groups[srcGroupID] = g
}
}
}
for _, dstGroupID := range rule.Destinations {
if g := a.Groups[dstGroupID]; g != nil {
if _, exists := components.Groups[dstGroupID]; !exists {
components.Groups[dstGroupID] = g.ToComponent()
components.Groups[dstGroupID] = g
}
}
}
@@ -276,22 +278,20 @@ func (a *Account) GetPeerNetworkMapComponents(
// network in the account — accounts with many tenants/networks
// shipped tens of unrelated peers in `peers[]` and `routers_map`.
if addSourcePeers {
components.RoutersMap[resource.NetworkID] = routerTypes.ToComponentMap(networkRoutingPeers)
components.RoutersMap[resource.NetworkID] = networkRoutingPeers
for peerIDKey := range networkRoutingPeers {
if p := a.Peers[peerIDKey]; p != nil {
cp := components.RouterPeers[peerIDKey]
if cp == nil {
cp = p.ToComponent()
components.RouterPeers[peerIDKey] = cp
if _, exists := components.RouterPeers[peerIDKey]; !exists {
components.RouterPeers[peerIDKey] = p
}
if _, exists := components.Peers[peerIDKey]; !exists {
if _, validated := validatedPeersMap[peerIDKey]; validated {
components.Peers[peerIDKey] = cp
components.Peers[peerIDKey] = p
}
}
}
}
components.NetworkResources = append(components.NetworkResources, resource.ToComponent())
components.NetworkResources = append(components.NetworkResources, resource)
}
}
@@ -312,14 +312,14 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
peerSSHEnabled bool,
validatedPeersMap map[string]struct{},
postureFailedPeers *map[string]map[string]struct{},
) (map[string]*ComponentPeer, map[string]*Group, []*Policy, []*route.Route, sshRequirements) {
relevantPeerIDs := make(map[string]*ComponentPeer, len(a.Peers)/4)
) (map[string]*nbpeer.Peer, map[string]*Group, []*Policy, []*route.Route, sshRequirements) {
relevantPeerIDs := make(map[string]*nbpeer.Peer, len(a.Peers)/4)
relevantGroupIDs := make(map[string]*Group, len(a.Groups)/4)
relevantPolicies := make([]*Policy, 0, len(a.Policies))
relevantRoutes := make([]*route.Route, 0, len(a.Routes))
sshReqs := sshRequirements{neededGroupIDs: make(map[string]struct{})}
relevantPeerIDs[peerID] = a.GetPeer(peerID).ToComponent()
relevantPeerIDs[peerID] = a.GetPeer(peerID)
peerGroupSet := make(map[string]struct{}, 8)
for groupID, group := range a.Groups {
@@ -384,7 +384,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
if r.Peer != "" {
if _, ok := validatedPeersMap[r.Peer]; ok {
if p := a.GetPeer(r.Peer); p != nil {
relevantPeerIDs[r.Peer] = p.ToComponent()
relevantPeerIDs[r.Peer] = p
}
}
}
@@ -401,7 +401,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
continue
}
if p := a.GetPeer(pid); p != nil {
relevantPeerIDs[pid] = p.ToComponent()
relevantPeerIDs[pid] = p
}
}
}
@@ -458,9 +458,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
if peerInSources {
policyRelevant = true
for _, pid := range destinationPeers {
if _, exists := relevantPeerIDs[pid]; !exists {
relevantPeerIDs[pid] = a.GetPeer(pid).ToComponent()
}
relevantPeerIDs[pid] = a.GetPeer(pid)
}
for _, dstGroupID := range rule.Destinations {
relevantGroupIDs[dstGroupID] = a.GetGroup(dstGroupID)
@@ -470,9 +468,7 @@ func (a *Account) getPeersGroupsPoliciesRoutes(
if peerInDestinations {
policyRelevant = true
for _, pid := range sourcePeers {
if _, exists := relevantPeerIDs[pid]; !exists {
relevantPeerIDs[pid] = a.GetPeer(pid).ToComponent()
}
relevantPeerIDs[pid] = a.GetPeer(pid)
}
for _, srcGroupID := range rule.Sources {
relevantGroupIDs[srcGroupID] = a.GetGroup(srcGroupID)
@@ -628,7 +624,7 @@ func (a *Account) getPostureValidPeersSaveFailed(inputPeers []string, postureChe
// that name them. Calculate() tolerates groups with empty Peers (the inner
// loops simply iterate zero times), so retaining them is behaviourally a
// no-op for the legacy path that consumes the same NetworkMapComponents.
func filterGroupPeers(groups *map[string]*ComponentGroup, peers map[string]*ComponentPeer) {
func filterGroupPeers(groups *map[string]*Group, peers map[string]*nbpeer.Peer) {
for groupID, groupInfo := range *groups {
filteredPeers := make([]string, 0, len(groupInfo.Peers))
for _, pid := range groupInfo.Peers {
@@ -638,14 +634,14 @@ func filterGroupPeers(groups *map[string]*ComponentGroup, peers map[string]*Comp
}
if len(filteredPeers) != len(groupInfo.Peers) {
ng := *groupInfo
ng := groupInfo.Copy()
ng.Peers = filteredPeers
(*groups)[groupID] = &ng
(*groups)[groupID] = ng
}
}
}
func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}, policies []*Policy, resourcePoliciesMap map[string][]*Policy, peers map[string]*ComponentPeer) {
func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}, policies []*Policy, resourcePoliciesMap map[string][]*Policy, peers map[string]*nbpeer.Peer) {
if len(*postureFailedPeers) == 0 {
return
}
@@ -680,7 +676,7 @@ func filterPostureFailedPeers(postureFailedPeers *map[string]map[string]struct{}
}
}
func filterDNSRecordsByPeers(records []nbdns.SimpleRecord, peers map[string]*ComponentPeer, includeIPv6 bool) []nbdns.SimpleRecord {
func filterDNSRecordsByPeers(records []nbdns.SimpleRecord, peers map[string]*nbpeer.Peer, includeIPv6 bool) []nbdns.SimpleRecord {
if len(records) == 0 || len(peers) == 0 {
return nil
}

View File

@@ -9,6 +9,7 @@ import (
"github.com/stretchr/testify/require"
nbdns "github.com/netbirdio/netbird/dns"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
)
func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
@@ -48,7 +49,7 @@ func TestPrivateService_NetworkMap_UserPeer_AndProxyPeer(t *testing.T) {
})
}
func netmapPeerIDs(peers []*ComponentPeer) []string {
func netmapPeerIDs(peers []*nbpeer.Peer) []string {
ids := make([]string, 0, len(peers))
for _, p := range peers {
if p == nil {

View File

@@ -666,7 +666,7 @@ func Test_ExpandPortsAndRanges_SSHRuleExpansion(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := ExpandPortsAndRanges(tt.base, tt.rule, tt.peer.ToComponent())
result := ExpandPortsAndRanges(tt.base, tt.rule, tt.peer)
var ports []string
for _, fr := range result {

View File

@@ -6,6 +6,7 @@ import (
"net"
"net/netip"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
nbroute "github.com/netbirdio/netbird/route"
sharedtypes "github.com/netbirdio/netbird/shared/management/types"
)
@@ -17,6 +18,9 @@ type DNSSettings = sharedtypes.DNSSettings
type FirewallRule = sharedtypes.FirewallRule
type Group = sharedtypes.Group
type GroupPeer = sharedtypes.GroupPeer
type Network = sharedtypes.Network
type NetworkMap = sharedtypes.NetworkMap
type ForwardingRule = sharedtypes.ForwardingRule
@@ -38,18 +42,6 @@ type RouteFirewallRule = sharedtypes.RouteFirewallRule
type NetworkMapComponents = sharedtypes.NetworkMapComponents
type ComponentPeer = sharedtypes.ComponentPeer
type ComponentGroup = sharedtypes.ComponentGroup
type ComponentRouter = sharedtypes.ComponentRouter
type ComponentResource = sharedtypes.ComponentResource
type ComponentResourceType = sharedtypes.ComponentResourceType
const (
ComponentResourceHost = sharedtypes.ComponentResourceHost
ComponentResourceSubnet = sharedtypes.ComponentResourceSubnet
ComponentResourceDomain = sharedtypes.ComponentResourceDomain
)
var EmptyNetworkMapComponents = sharedtypes.EmptyNetworkMapComponents
type AccountSettingsInfo = sharedtypes.AccountSettingsInfo
@@ -60,7 +52,12 @@ type NetworkMapComponentsCompact = sharedtypes.NetworkMapComponentsCompact
type LookupMap = sharedtypes.LookupMap
type FirewallRuleContext = sharedtypes.FirewallRuleContext
const GroupAllName = sharedtypes.GroupAllName
const (
GroupIssuedAPI = sharedtypes.GroupIssuedAPI
GroupIssuedJWT = sharedtypes.GroupIssuedJWT
GroupIssuedIntegration = sharedtypes.GroupIssuedIntegration
GroupAllName = sharedtypes.GroupAllName
)
// Function forwarders preserve types.X(...) call sites that previously
// resolved to package-local funcs. Plain forwarders (not var aliases) keep
@@ -70,11 +67,11 @@ func PolicyRuleImpliesLegacySSH(rule *PolicyRule) bool {
return sharedtypes.PolicyRuleImpliesLegacySSH(rule)
}
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *ComponentPeer) []*FirewallRule {
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *nbpeer.Peer) []*FirewallRule {
return sharedtypes.ExpandPortsAndRanges(base, rule, peer)
}
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *ComponentPeer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *nbpeer.Peer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
return sharedtypes.AppendIPv6FirewallRule(rules, rulesExists, peer, targetPeer, rule, rc)
}
@@ -82,7 +79,7 @@ func CalculateNetworkMapFromComponents(ctx context.Context, components *NetworkM
return sharedtypes.CalculateNetworkMapFromComponents(ctx, components)
}
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*ComponentPeer, direction int, includeIPv6 bool) []*RouteFirewallRule {
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int, includeIPv6 bool) []*RouteFirewallRule {
return sharedtypes.GenerateRouteFirewallRules(ctx, route, rule, groupPeers, direction, includeIPv6)
}

View File

@@ -9,7 +9,6 @@ import (
"github.com/stretchr/testify/require"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
)
func TestNetworkMapComponents_IPv6EndToEnd(t *testing.T) {
@@ -105,7 +104,7 @@ func TestNetworkMapComponents_RemotePeerWithoutCapability(t *testing.T) {
require.NotNil(t, nm)
t.Run("AllowedIPs include remote v6", func(t *testing.T) {
var dst *types.ComponentPeer
var dst *nbpeer.Peer
for _, p := range nm.Peers {
if p.ID == "peer-dst-1" {
dst = p

View File

@@ -49,7 +49,7 @@ func allPeersValidated(account *types.Account, excludePeerIDs ...string) map[str
return validated
}
func peerIDs(peers []*types.ComponentPeer) []string {
func peerIDs(peers []*nbpeer.Peer) []string {
ids := make([]string, len(peers))
for i, p := range peers {
ids[i] = p.ID

View File

@@ -19,3 +19,34 @@ func Difference(a, b []string) []string {
func ToPtr[T any](value T) *T {
return &value
}
type comparableObject[T any] interface {
Equal(other T) bool
}
func MergeUnique[T comparableObject[T]](arr1, arr2 []T) []T {
var result []T
for _, item := range arr1 {
if !contains(result, item) {
result = append(result, item)
}
}
for _, item := range arr2 {
if !contains(result, item) {
result = append(result, item)
}
}
return result
}
func contains[T comparableObject[T]](slice []T, element T) bool {
for _, item := range slice {
if item.Equal(element) {
return true
}
}
return false
}

View File

@@ -1,4 +1,4 @@
package types
package util
import (
"testing"
@@ -17,7 +17,7 @@ func (t testObject) Equal(other testObject) bool {
func Test_MergeUniqueArraysWithoutDuplicates(t *testing.T) {
arr1 := []testObject{{value: 1}, {value: 2}}
arr2 := []testObject{{value: 2}, {value: 3}}
result := mergeUnique(arr1, arr2)
result := MergeUnique(arr1, arr2)
assert.Len(t, result, 3)
assert.Contains(t, result, testObject{value: 1})
assert.Contains(t, result, testObject{value: 2})
@@ -27,14 +27,14 @@ func Test_MergeUniqueArraysWithoutDuplicates(t *testing.T) {
func Test_MergeUniqueHandlesEmptyArrays(t *testing.T) {
arr1 := []testObject{}
arr2 := []testObject{}
result := mergeUnique(arr1, arr2)
result := MergeUnique(arr1, arr2)
assert.Empty(t, result)
}
func Test_MergeUniqueHandlesOneEmptyArray(t *testing.T) {
arr1 := []testObject{{value: 1}, {value: 2}}
arr2 := []testObject{}
result := mergeUnique(arr1, arr2)
result := MergeUnique(arr1, arr2)
assert.Len(t, result, 2)
assert.Contains(t, result, testObject{value: 1})
assert.Contains(t, result, testObject{value: 2})

View File

@@ -221,21 +221,14 @@ func (l *Logger) allowDenyLog(serviceID types.ServiceID, reason string) bool {
// proxy/internal/middleware/keys.go — only the dimensions management needs to
// record a usage row (provider / model / tokens / cost / groups).
var usageMetadataKeys = map[string]struct{}{
"llm.provider": {},
"llm.model": {},
"llm.resolved_provider_id": {},
"llm.input_tokens": {},
"llm.output_tokens": {},
"llm.total_tokens": {},
"llm.cached_input_tokens": {},
"llm.cache_creation_tokens": {},
"cost.usd_input": {},
"cost.usd_cached_input": {},
"cost.usd_cache_creation": {},
"cost.usd_output": {},
"cost.usd_total": {},
"cost.usd_cache": {},
"llm.authorising_groups": {},
"llm.provider": {},
"llm.model": {},
"llm.resolved_provider_id": {},
"llm.input_tokens": {},
"llm.output_tokens": {},
"llm.total_tokens": {},
"cost.usd_total": {},
"llm.authorising_groups": {},
}
// stripAgentNetworkEntryForUsage returns the entry reduced to what's needed to

View File

@@ -56,12 +56,10 @@ type bedrockResponse struct {
OutputTokens int64 `json:"output_tokens"`
CacheReadInputTokens int64 `json:"cache_read_input_tokens"`
CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"`
// Converse — camelCase; cache buckets are additive to inputTokens (AWS names the write bucket cacheWriteInputTokens).
InputTokensCamel int64 `json:"inputTokens"`
OutputTokensCamel int64 `json:"outputTokens"`
TotalTokensCamel int64 `json:"totalTokens"`
CacheReadTokensCamel int64 `json:"cacheReadInputTokens"`
CacheWriteTokensCamel int64 `json:"cacheWriteInputTokens"`
// Converse — camelCase.
InputTokensCamel int64 `json:"inputTokens"`
OutputTokensCamel int64 `json:"outputTokens"`
TotalTokensCamel int64 `json:"totalTokens"`
} `json:"usage"`
}
@@ -85,18 +83,16 @@ func (BedrockParser) ParseResponse(status int, contentType string, body []byte)
}
inTok := firstNonZero(resp.Usage.InputTokens, resp.Usage.InputTokensCamel)
outTok := firstNonZero(resp.Usage.OutputTokens, resp.Usage.OutputTokensCamel)
cacheRead := firstNonZero(resp.Usage.CacheReadInputTokens, resp.Usage.CacheReadTokensCamel)
cacheWrite := firstNonZero(resp.Usage.CacheCreationInputTokens, resp.Usage.CacheWriteTokensCamel)
total := resp.Usage.TotalTokensCamel
if total == 0 {
total = inTok + outTok + cacheRead + cacheWrite
total = inTok + outTok + resp.Usage.CacheReadInputTokens + resp.Usage.CacheCreationInputTokens
}
return Usage{
InputTokens: inTok,
OutputTokens: outTok,
TotalTokens: total,
CachedInputTokens: cacheRead,
CacheCreationTokens: cacheWrite,
CachedInputTokens: resp.Usage.CacheReadInputTokens,
CacheCreationTokens: resp.Usage.CacheCreationInputTokens,
}, nil
}

View File

@@ -26,18 +26,6 @@ func TestBedrockParser_ParseResponse_Converse(t *testing.T) {
require.Equal(t, int64(14), u.TotalTokens, "converse uses provider total")
}
// Converse camelCase cache fields must land in the billed Usage buckets, same as the InvokeModel snake_case fields.
func TestBedrockParser_ParseResponse_ConverseCacheBuckets(t *testing.T) {
body := []byte(`{"usage":{"inputTokens":11,"outputTokens":3,"cacheReadInputTokens":7,"cacheWriteInputTokens":9}}`)
u, err := BedrockParser{}.ParseResponse(200, "application/json", body)
require.NoError(t, err)
require.Equal(t, int64(11), u.InputTokens, "converse input tokens")
require.Equal(t, int64(3), u.OutputTokens, "converse output tokens")
require.Equal(t, int64(7), u.CachedInputTokens, "converse cache-read tokens")
require.Equal(t, int64(9), u.CacheCreationTokens, "converse cache-write tokens")
require.Equal(t, int64(11+3+7+9), u.TotalTokens, "total backfill is additive when the provider omits totalTokens")
}
func TestBedrockParser_ParseResponse_StreamingUnsupported(t *testing.T) {
_, err := BedrockParser{}.ParseResponse(200, "application/vnd.amazon.eventstream", []byte("binary"))
require.ErrorIs(t, err, ErrStreamingUnsupported, "event-stream must route to the streaming accumulator")

View File

@@ -128,46 +128,6 @@ type Table struct {
// - Other providers: cached and cacheCreation are ignored; cost is
// inTokens*InputPer1K + outTokens*OutputPer1K.
func (t *Table) Cost(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (float64, bool) {
c, ok := t.Costs(provider, model, inTokens, outTokens, cachedInput, cacheCreation)
return c.TotalUSD, ok
}
// Costs is a per-request cost split. The four per-bucket fields are the base
// of the breakdown — one per token bucket the provider bills separately — and
// the two aggregates are derived from them:
//
// TotalUSD = InputUSD + CachedInputUSD + CacheCreationUSD + OutputUSD
// CacheUSD = CachedInputUSD + CacheCreationUSD
//
// InputUSD is always the cost of the *non-cached* input bucket, for both
// provider shapes: on OpenAI the cached subset is carved out of inTokens and
// billed as CachedInputUSD, so the two never double-count. Buckets a provider
// doesn't bill are zero, which keeps the identities above true everywhere.
type Costs struct {
InputUSD float64
CachedInputUSD float64
CacheCreationUSD float64
OutputUSD float64
TotalUSD float64
CacheUSD float64
}
// newCosts assembles a split from its per-bucket parts, deriving the two
// aggregates so TotalUSD and CacheUSD can never drift from the breakdown.
func newCosts(input, cachedInput, cacheCreation, output float64) Costs {
return Costs{
InputUSD: input,
CachedInputUSD: cachedInput,
CacheCreationUSD: cacheCreation,
OutputUSD: output,
TotalUSD: input + cachedInput + cacheCreation + output,
CacheUSD: cachedInput + cacheCreation,
}
}
// Costs returns the estimated USD cost split for the given token counts, with
// the same semantics as Cost.
func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput, cacheCreation int64) (Costs, bool) {
// Clamp negatives to zero before any pricing math so a malformed
// upstream count can never produce a negative cost.
if inTokens < 0 {
@@ -183,15 +143,15 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput,
cacheCreation = 0
}
if t == nil {
return Costs{}, false
return 0, false
}
byModel, ok := t.entries[provider]
if !ok {
return Costs{}, false
return 0, false
}
entry, ok := byModel[model]
if !ok {
return Costs{}, false
return 0, false
}
output := (float64(outTokens) / 1000.0) * entry.OutputPer1K
switch provider {
@@ -208,7 +168,7 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput,
}
nonCached := float64(inTokens-clamped) / 1000.0 * entry.InputPer1K
cached := float64(clamped) / 1000.0 * cachedRate
return newCosts(nonCached, cached, 0, output), true
return nonCached + cached + output, true
case "anthropic", "bedrock":
// Bedrock-Anthropic returns the same additive cache buckets as
// first-party Anthropic; non-Anthropic Bedrock models simply report
@@ -224,10 +184,10 @@ func (t *Table) Costs(provider, model string, inTokens, outTokens, cachedInput,
input := float64(inTokens) / 1000.0 * entry.InputPer1K
read := float64(cachedInput) / 1000.0 * readRate
create := float64(cacheCreation) / 1000.0 * createRate
return newCosts(input, read, create, output), true
return input + read + create + output, true
default:
input := float64(inTokens) / 1000.0 * entry.InputPer1K
return newCosts(input, 0, 0, output), true
return input + output, true
}
}

View File

@@ -1,329 +0,0 @@
package builtin_test
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"strconv"
"testing"
"github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/proxy/internal/middleware"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/cost_meter"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_request_parser"
"github.com/netbirdio/netbird/proxy/internal/middleware/builtin/llm_response_parser"
)
// Drives the real pipeline (llm_request_parser → llm_response_parser → cost_meter) on the embedded default pricing
// table and asserts exact USD amounts hardcoded from the vendors' published prices, including the cache split.
func TestCostCalculation_ProviderMatrix(t *testing.T) {
// Empty data dir → embedded defaults, like a proxy with no pricing override.
builtin.Configure(context.Background(), t.TempDir(), nil, nil, nil)
reqMW, err := llm_request_parser.Factory{}.New(nil)
require.NoError(t, err, "build llm_request_parser")
respMW, err := llm_response_parser.Factory{}.New(nil)
require.NoError(t, err, "build llm_response_parser")
costMW, err := cost_meter.Factory{}.New(nil)
require.NoError(t, err, "build cost_meter")
t.Cleanup(func() { _ = costMW.Close() })
const jsonCT = "application/json"
const sseCT = "text/event-stream"
const awsCT = "application/vnd.amazon.eventstream"
cases := []struct {
name string
url string
reqBody []byte
respCT string
respBody []byte
wantProvider string
wantModel string
wantCost float64 // exact expected USD; ignored when wantSkip is set
wantCacheCost float64 // expected cost.usd_cache portion of wantCost
wantSkip string // expected cost.skipped reason, "" when priced
}{
{
// gpt-4o-mini $0.15/$0.60 per MTok: 1000×0.15/1M + 500×0.60/1M.
name: "openai chat completions",
url: "https://api.openai.com/v1/chat/completions",
reqBody: []byte(`{"model":"gpt-4o-mini","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"choices":[{"message":{"content":"pong"}}],"usage":{"prompt_tokens":1000,"completion_tokens":500,"total_tokens":1500}}`),
wantProvider: "openai",
wantModel: "gpt-4o-mini",
wantCost: 0.00045,
},
{
// OpenAI cached tokens are a SUBSET of prompt_tokens at a discount; gpt-4o $2.50/$10 per MTok, cached $1.25/M:
// 250×2.5/1M + 750×1.25/1M + 500×10/1M.
name: "openai cached subset discount",
url: "https://api.openai.com/v1/chat/completions",
reqBody: []byte(`{"model":"gpt-4o","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"prompt_tokens":1000,"completion_tokens":500,"prompt_tokens_details":{"cached_tokens":750}}}`),
wantProvider: "openai",
wantModel: "gpt-4o",
wantCost: 0.0065625,
wantCacheCost: 0.0009375,
},
{
// OpenAI streaming: usage rides the final SSE frame.
name: "openai chat SSE stream",
url: "https://api.openai.com/v1/chat/completions",
reqBody: []byte(`{"model":"gpt-4o-mini","stream":true,"messages":[{"role":"user","content":"hi"}]}`),
respCT: sseCT,
respBody: sseBody(`{"choices":[{"delta":{"content":"po"}}]}`, `{"choices":[{"delta":{"content":"ng"}}]}`, `{"choices":[],"usage":{"prompt_tokens":1000,"completion_tokens":500}}`, "[DONE]"),
wantProvider: "openai",
wantModel: "gpt-4o-mini",
wantCost: 0.00045,
},
{
// Mistral speaks the OpenAI shape: mistral-large-latest $0.50/$1.50 per MTok.
name: "mistral via openai shape",
url: "https://api.mistral.ai/v1/chat/completions",
reqBody: []byte(`{"model":"mistral-large-latest","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"prompt_tokens":1000,"completion_tokens":1000}}`),
wantProvider: "openai",
wantModel: "mistral-large-latest",
wantCost: 0.002,
},
{
// The field report, minus caching: Bedrock Sonnet 4.6 $3/$15 per MTok, 3×3/1M + 1514×15/1M = $0.022719.
// Also covers inference-profile normalization of the region-prefixed versioned id in the URL.
name: "bedrock invoke — reported scenario, no cache",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/global.anthropic.claude-sonnet-4-6-20260115-v1:0/invoke",
reqBody: []byte(`{"messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"content":[{"type":"text","text":"pong"}],"usage":{"input_tokens":3,"output_tokens":1514}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-sonnet-4-6",
wantCost: 0.022719,
},
{
// The field report as observed: the FIRST call of a session also wrote a 30,528-token prompt cache at
// 1.25× input ($3.75/M): 0.022719 + 30528×3.75/1M = $0.137199 — the reported $0.1372.
name: "bedrock invoke — reported scenario with cache write",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/global.anthropic.claude-sonnet-4-6-20260115-v1:0/invoke",
reqBody: []byte(`{"messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"input_tokens":3,"output_tokens":1514,"cache_creation_input_tokens":30528,"cache_read_input_tokens":0}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-sonnet-4-6",
wantCost: 0.137199,
wantCacheCost: 0.11448,
},
{
// Same numbers over the InvokeModel event-stream: message_start carries input + cache, message_delta the output.
name: "bedrock invoke stream with cache write",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/global.anthropic.claude-sonnet-4-6-20260115-v1:0/invoke-with-response-stream",
reqBody: []byte(`{"messages":[{"role":"user","content":"hi"}]}`),
respCT: awsCT,
respBody: bedrockInvokeStream(t, `{"type":"message_start","message":{"usage":{"input_tokens":3,"output_tokens":1,"cache_creation_input_tokens":30528}}}`, `{"type":"content_block_delta","delta":{"type":"text_delta","text":"pong"}}`, `{"type":"message_delta","usage":{"output_tokens":1514}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-sonnet-4-6",
wantCost: 0.137199,
wantCacheCost: 0.11448,
},
{
// Converse camelCase usage incl. cache buckets. Haiku 4.5 $1/$5 per MTok, read $0.10/M, write $1.25/M:
// 50×1/1M + 100×5/1M + 2000×0.1/1M + 1000×1.25/1M = $0.002.
name: "bedrock converse with cache buckets",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/eu.anthropic.claude-haiku-4-5-20251001-v1:0/converse",
reqBody: []byte(`{"messages":[{"role":"user","content":[{"text":"hi"}]}]}`),
respCT: jsonCT,
respBody: []byte(`{"output":{"message":{"content":[{"text":"pong"}]}},"usage":{"inputTokens":50,"outputTokens":100,"totalTokens":3150,"cacheReadInputTokens":2000,"cacheWriteInputTokens":1000}}`),
wantProvider: "bedrock",
wantModel: "anthropic.claude-haiku-4-5",
wantCost: 0.002,
wantCacheCost: 0.00145,
},
{
// Same numbers over converse-stream: usage rides the trailing metadata frame.
name: "bedrock converse stream with cache buckets",
url: "https://bedrock-runtime.eu-central-1.amazonaws.com/model/eu.anthropic.claude-haiku-4-5-20251001-v1:0/converse-stream",
reqBody: []byte(`{"messages":[{"role":"user","content":[{"text":"hi"}]}]}`),
respCT: awsCT,
respBody: bedrockConverseStream(t,
`{"delta":{"text":"pong"}}`,
`{"usage":{"inputTokens":50,"outputTokens":100,"totalTokens":3150,"cacheReadInputTokens":2000,"cacheWriteInputTokens":1000}}`,
),
wantProvider: "bedrock",
wantModel: "anthropic.claude-haiku-4-5",
wantCost: 0.002,
wantCacheCost: 0.00145,
},
{
// First-party Anthropic, additive cache buckets. Sonnet 4.6:
// 256×3/1M + 200×15/1M + 768×0.3/1M + 512×3.75/1M.
name: "anthropic messages with cache buckets",
url: "https://api.anthropic.com/v1/messages",
reqBody: []byte(`{"model":"claude-sonnet-4-6","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"content":[{"type":"text","text":"pong"}],"usage":{"input_tokens":256,"output_tokens":200,"cache_read_input_tokens":768,"cache_creation_input_tokens":512}}`),
wantProvider: "anthropic",
wantModel: "claude-sonnet-4-6",
wantCost: 0.0059184,
wantCacheCost: 0.0021504,
},
{
// Anthropic SSE: input from message_start, output from message_delta. Haiku 4.5: 1000×1/1M + 2000×5/1M.
name: "anthropic SSE stream",
url: "https://api.anthropic.com/v1/messages",
reqBody: []byte(`{"model":"claude-haiku-4-5","stream":true,"messages":[{"role":"user","content":"hi"}]}`),
respCT: sseCT,
respBody: sseBody(`{"type":"message_start","message":{"usage":{"input_tokens":1000,"output_tokens":2}}}`, `{"type":"content_block_delta","delta":{"type":"text_delta","text":"pong"}}`, `{"type":"message_delta","usage":{"output_tokens":2000}}`, `{"type":"message_stop"}`),
wantProvider: "anthropic",
wantModel: "claude-haiku-4-5",
wantCost: 0.011,
},
{
// Kimi's Anthropic-compatible endpoint: kimi-k3 $3/$15 per MTok under the anthropic table.
name: "kimi anthropic shape",
url: "https://api.moonshot.ai/anthropic/v1/messages",
reqBody: []byte(`{"model":"kimi-k3","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"input_tokens":1000,"output_tokens":1000}}`),
wantProvider: "anthropic",
wantModel: "kimi-k3",
wantCost: 0.018,
},
{
// Vertex path-routed model with "@version" stripped; Anthropic-on-Vertex priced under the anthropic table.
name: "vertex anthropic path-routed",
url: "https://aiplatform.googleapis.com/v1/projects/p/locations/global/publishers/anthropic/models/claude-sonnet-4-6@20260115:rawPredict",
reqBody: []byte(`{"anthropic_version":"vertex-2023-10-16","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"input_tokens":200,"output_tokens":100}}`),
wantProvider: "anthropic",
wantModel: "claude-sonnet-4-6",
wantCost: 0.0021,
},
{
// Gateway-prefixed model ids are not in the pricing table: the meter must SKIP, never guess a rate.
name: "gateway-prefixed model skips pricing",
url: "https://gateway.example.com/v1/chat/completions",
reqBody: []byte(`{"model":"openai/gpt-4o-mini","messages":[{"role":"user","content":"hi"}]}`),
respCT: jsonCT,
respBody: []byte(`{"usage":{"prompt_tokens":1000,"completion_tokens":500}}`),
wantProvider: "openai",
wantModel: "openai/gpt-4o-mini",
wantSkip: "unknown_model",
},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
in := &middleware.Input{
Method: "POST",
URL: tc.url,
Headers: []middleware.KV{{Key: "Content-Type", Value: "application/json"}},
Body: tc.reqBody,
}
reqOut, err := reqMW.Invoke(context.Background(), in)
require.NoError(t, err, "request parser")
in.Metadata = append(in.Metadata, reqOut.Metadata...)
require.Equal(t, tc.wantProvider, metaKV(in.Metadata, middleware.KeyLLMProvider), "detected provider")
require.Equal(t, tc.wantModel, metaKV(in.Metadata, middleware.KeyLLMModel), "detected (normalized) model")
in.Status = 200
in.RespHeaders = []middleware.KV{{Key: "Content-Type", Value: tc.respCT}}
in.RespBody = tc.respBody
respOut, err := respMW.Invoke(context.Background(), in)
require.NoError(t, err, "response parser")
in.Metadata = append(in.Metadata, respOut.Metadata...)
costOut, err := costMW.Invoke(context.Background(), in)
require.NoError(t, err, "cost meter")
if tc.wantSkip != "" {
assert.Equal(t, tc.wantSkip, metaKV(costOut.Metadata, middleware.KeyCostSkipped), "expected cost skip reason")
assert.Empty(t, metaKV(costOut.Metadata, middleware.KeyCostUSDTotal), "no cost may be emitted on skip")
return
}
raw := metaKV(costOut.Metadata, middleware.KeyCostUSDTotal)
require.NotEmpty(t, raw, "cost.usd_total must be emitted; skip=%q", metaKV(costOut.Metadata, middleware.KeyCostSkipped))
got, err := strconv.ParseFloat(raw, 64)
require.NoError(t, err, "cost must be a float")
// cost.usd_total is rendered with %.6f: allow half of the last printed digit on top of float error.
assert.InDelta(t, tc.wantCost, got, 5.1e-7, "USD cost for %s", tc.name)
rawCache := metaKV(costOut.Metadata, middleware.KeyCostUSDCache)
require.NotEmpty(t, rawCache, "cost.usd_cache must be emitted next to cost.usd_total")
gotCache, err := strconv.ParseFloat(rawCache, 64)
require.NoError(t, err, "cache cost must be a float")
assert.InDelta(t, tc.wantCacheCost, gotCache, 5.1e-7, "cache USD cost for %s", tc.name)
})
}
}
// metaKV returns the value for key in kvs, or "" when absent.
func metaKV(kvs []middleware.KV, key string) string {
for _, kv := range kvs {
if kv.Key == key {
return kv.Value
}
}
return ""
}
// sseBody renders data frames as a text/event-stream body.
func sseBody(frames ...string) []byte {
var b bytes.Buffer
for _, f := range frames {
b.WriteString("data: ")
b.WriteString(f)
b.WriteString("\n\n")
}
return b.Bytes()
}
// awsFrame encodes one AWS event-stream frame with the given :event-type.
func awsFrame(t *testing.T, eventType string, payload []byte) []byte {
t.Helper()
var buf bytes.Buffer
enc := eventstream.NewEncoder()
require.NoError(t, enc.Encode(&buf, eventstream.Message{
Headers: eventstream.Headers{{Name: ":event-type", Value: eventstream.StringValue(eventType)}},
Payload: payload,
}), "encode event-stream frame")
return buf.Bytes()
}
// bedrockInvokeStream builds an invoke-with-response-stream body: each "chunk" frame wraps a base64 Anthropic event.
func bedrockInvokeStream(t *testing.T, events ...string) []byte {
t.Helper()
var body bytes.Buffer
for _, ev := range events {
wrap, err := json.Marshal(map[string]string{"bytes": base64.StdEncoding.EncodeToString([]byte(ev))})
require.NoError(t, err)
body.Write(awsFrame(t, "chunk", wrap))
}
return body.Bytes()
}
// bedrockConverseStream builds a converse-stream body: contentBlockDelta frames plus a trailing metadata usage frame.
func bedrockConverseStream(t *testing.T, deltas ...string) []byte {
t.Helper()
var body bytes.Buffer
for i, ev := range deltas {
eventType := "contentBlockDelta"
if i == len(deltas)-1 {
eventType = "metadata"
}
body.Write(awsFrame(t, eventType, []byte(ev)))
}
return body.Bytes()
}

View File

@@ -32,12 +32,7 @@ const (
)
var metadataKeys = []string{
middleware.KeyCostUSDInput,
middleware.KeyCostUSDCachedInput,
middleware.KeyCostUSDCacheCreation,
middleware.KeyCostUSDOutput,
middleware.KeyCostUSDTotal,
middleware.KeyCostUSDCache,
middleware.KeyCostSkipped,
}
@@ -145,38 +140,18 @@ func (m *Middleware) Invoke(_ context.Context, in *middleware.Input) (*middlewar
}
table := m.loader.Get()
costs, ok := table.Costs(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
cost, ok := table.Cost(provider, model, inTokens, outTokens, cachedTokens, cacheCreationTokens)
if !ok {
out.Metadata = skip(skipUnknownModel)
return out, nil
}
// Per-bucket costs first: they're the base of the breakdown, and the two
// aggregates that follow are derived from exactly these four values.
out.Metadata = []middleware.KV{
{Key: middleware.KeyCostUSDInput, Value: usd(costs.InputUSD)},
{Key: middleware.KeyCostUSDCachedInput, Value: usd(costs.CachedInputUSD)},
{Key: middleware.KeyCostUSDCacheCreation, Value: usd(costs.CacheCreationUSD)},
{Key: middleware.KeyCostUSDOutput, Value: usd(costs.OutputUSD)},
{Key: middleware.KeyCostUSDTotal, Value: usd(costs.TotalUSD)},
{Key: middleware.KeyCostUSDCache, Value: usd(costs.CacheUSD)},
{Key: middleware.KeyCostUSDTotal, Value: fmt.Sprintf("%.6f", cost)},
}
return out, nil
}
// usd renders a cost as the fixed-precision string every cost.usd_* key
// carries, so the per-bucket values and the aggregates round identically.
//
// 9 decimals, not 6: these values are summed downstream — per request, per
// session, and per usage bucket — so the rounding step is applied once per
// bucket per row and then accumulated. At 6 decimals a single row loses up to
// 2e-6 across its four buckets (enough to break a 1e-6 reconciliation against
// published rates), and a bucket smaller than half a microdollar quantises to
// zero outright: 16 cache-read tokens on a cheap model is 1.6e-9, so summing
// 10k such rows reports 0.02 instead of 0.016. Nano-dollar precision keeps the
// per-row error ~1000x below the smallest realistic bucket.
func usd(v float64) string { return fmt.Sprintf("%.9f", v) }
// skip returns a single-entry metadata slice carrying the given skip
// reason under KeyCostSkipped.
func skip(reason string) []middleware.KV {

View File

@@ -67,15 +67,7 @@ func TestMiddleware_StaticSurface(t *testing.T) {
assert.NoError(t, mw.Close(), "Close on stateless middleware is a no-op")
keys := mw.MetadataKeys()
expected := []string{
middleware.KeyCostUSDInput,
middleware.KeyCostUSDCachedInput,
middleware.KeyCostUSDCacheCreation,
middleware.KeyCostUSDOutput,
middleware.KeyCostUSDTotal,
middleware.KeyCostUSDCache,
middleware.KeyCostSkipped,
}
expected := []string{middleware.KeyCostUSDTotal, middleware.KeyCostSkipped}
assert.Equal(t, expected, keys, "metadata key allowlist must match the spec")
}
@@ -113,7 +105,7 @@ func TestFactory_DefaultPricingPathLoadsFixture(t *testing.T) {
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
require.True(t, ok, "cost.usd_total must be emitted for known model")
assert.Equal(t, "0.000750000", value, "0.00015 + 0.0006 per 1k tokens, 9-decimal format")
assert.Equal(t, "0.000750", value, "0.00015 + 0.0006 per 1k tokens, 6-decimal format")
}
func TestFactory_PricingPathOverride(t *testing.T) {
@@ -137,7 +129,7 @@ func TestFactory_PricingPathOverride(t *testing.T) {
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
require.True(t, ok, "cost.usd_total must be emitted with custom pricing path")
assert.Equal(t, "0.015000000", value, "2*0.0025 + 1*0.01 = 0.015 with 9-decimal format")
assert.Equal(t, "0.015000", value, "2*0.0025 + 1*0.01 = 0.015 with 6-decimal format")
}
func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
@@ -156,7 +148,7 @@ func TestInvoke_ComputesCostForKnownModel(t *testing.T) {
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
require.True(t, ok, "cost.usd_total must be emitted")
assert.Equal(t, "0.018000000", value, "0.003 + 0.015 = 0.018 with 9-decimal format")
assert.Equal(t, "0.018000", value, "0.003 + 0.015 = 0.018 with 6-decimal format")
_, skipped := metaValue(t, out.Metadata, middleware.KeyCostSkipped)
assert.False(t, skipped, "cost.skipped must not be set when cost is computed")
}
@@ -365,25 +357,8 @@ func TestInvoke_OpenAICachedSubsetDiscount(t *testing.T) {
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
require.True(t, ok, "cached subset path must produce a cost — never a skip")
// 250 non-cached at 0.0025/1k + 750 cached at 0.00125/1k + 500 output at 0.01/1k.
assert.Equal(t, "0.006562500", value,
assert.Equal(t, "0.006563", value,
"cached subset must be billed at the discount rate, non-cached at the full rate; never double-billed")
cache, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDCache)
require.True(t, ok, "cost.usd_cache must be emitted alongside cost.usd_total")
// 750 cached at 0.00125/1k = 0.0009375.
assert.Equal(t, "0.000937500", cache, "cache cost is the discounted cost of the cached subset")
// Per-bucket breakdown. On OpenAI the cached subset is carved out of the
// input bucket, so input covers only the 250 non-cached tokens — the two
// must never double-count the same 750 tokens.
assertBucket(t, out.Metadata, middleware.KeyCostUSDInput, "0.000625000",
"input bucket bills only the non-cached remainder")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCachedInput, "0.000937500",
"cached-input bucket bills the discounted subset")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCacheCreation, "0.000000000",
"OpenAI has no cache-write bucket")
assertBucket(t, out.Metadata, middleware.KeyCostUSDOutput, "0.005000000",
"output bucket bills 500 tokens at 0.01/1k")
}
// TestInvoke_AnthropicCacheBucketsAdditive proves the Anthropic
@@ -409,33 +384,9 @@ func TestInvoke_AnthropicCacheBucketsAdditive(t *testing.T) {
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
require.True(t, ok)
// 256 input * 0.003 + 768 cache_read * 0.0003 + 512 cache_creation * 0.00375 + 200 output * 0.015
// = 0.000768 + 0.0002304 + 0.00192 + 0.003 = 0.0059184.
assert.Equal(t, "0.005918400", value,
// = 0.000768 + 0.0002304 + 0.00192 + 0.003 = 0.0059184 → "0.005918" with 6-decimal format.
assert.Equal(t, "0.005918", value,
"each Anthropic input bucket must bill at its own rate — cache_read cheap, cache_creation expensive, regular input mid")
cache, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDCache)
require.True(t, ok, "cost.usd_cache must be emitted alongside cost.usd_total")
// 768 cache_read * 0.0003 + 512 cache_creation * 0.00375 = 0.0021504.
assert.Equal(t, "0.002150400", cache, "cache cost sums the read and creation buckets")
// Per-bucket breakdown: four separately-billed buckets, each at its own rate.
assertBucket(t, out.Metadata, middleware.KeyCostUSDInput, "0.000768000",
"input bucket bills 256 tokens at 0.003/1k")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCachedInput, "0.000230400",
"cache-read bucket bills 768 tokens at the cheap 0.0003/1k")
assertBucket(t, out.Metadata, middleware.KeyCostUSDCacheCreation, "0.001920000",
"cache-write bucket bills 512 tokens at the expensive 0.00375/1k")
assertBucket(t, out.Metadata, middleware.KeyCostUSDOutput, "0.003000000",
"output bucket bills 200 tokens at 0.015/1k")
}
// assertBucket asserts one per-bucket cost key carries the expected
// 6-decimal value.
func assertBucket(t *testing.T, md []middleware.KV, key, want, msg string) {
t.Helper()
got, ok := metaValue(t, md, key)
require.Truef(t, ok, "%s must be emitted", key)
assert.Equal(t, want, got, msg)
}
// TestInvoke_CachedTokensAbsentFallsBackToBaseFormula covers the
@@ -460,7 +411,7 @@ func TestInvoke_CachedTokensAbsentFallsBackToBaseFormula(t *testing.T) {
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
require.True(t, ok)
// 1000 input * 0.0025 + 500 output * 0.01 = 0.0025 + 0.005 = 0.0075
assert.Equal(t, "0.007500000", value, "no cached metadata = same cost as before the feature landed")
assert.Equal(t, "0.007500", value, "no cached metadata = same cost as before the feature landed")
}
// TestInvoke_UnparseableCachedTokensSkippedSilently proves the
@@ -484,7 +435,7 @@ func TestInvoke_UnparseableCachedTokensSkippedSilently(t *testing.T) {
require.NoError(t, err)
value, ok := metaValue(t, out.Metadata, middleware.KeyCostUSDTotal)
require.True(t, ok, "garbage cache metadata must NOT switch the response from a cost to a skip — fall back to 0 cached")
assert.Equal(t, "0.007500000", value, "same as the no-cached-metadata path")
assert.Equal(t, "0.007500", value, "same as the no-cached-metadata path")
}
// TestMiddleware_CloseCancelsReloader proves Close stops the per-instance

View File

@@ -69,18 +69,15 @@ func applyBedrockInvokeChunk(payload []byte, usage *llm.Usage, completion *strin
}
// converseStreamEvent captures the Converse stream frames carrying completion
// text (contentBlockDelta) and the final token usage (metadata). Cache buckets
// are additive to inputTokens (AWS write bucket: cacheWriteInputTokens).
// text (contentBlockDelta) and the final token usage (metadata).
type converseStreamEvent struct {
Delta *struct {
Text string `json:"text"`
} `json:"delta"`
Usage *struct {
InputTokens int64 `json:"inputTokens"`
OutputTokens int64 `json:"outputTokens"`
TotalTokens int64 `json:"totalTokens"`
CacheReadTokens int64 `json:"cacheReadInputTokens"`
CacheWriteTokens int64 `json:"cacheWriteInputTokens"`
InputTokens int64 `json:"inputTokens"`
OutputTokens int64 `json:"outputTokens"`
TotalTokens int64 `json:"totalTokens"`
} `json:"usage"`
}
@@ -108,12 +105,6 @@ func applyConverseStreamEvent(eventType string, payload []byte, usage *llm.Usage
if ev.Usage.TotalTokens > 0 {
usage.TotalTokens = ev.Usage.TotalTokens
}
if ev.Usage.CacheReadTokens > 0 {
usage.CachedInputTokens = ev.Usage.CacheReadTokens
}
if ev.Usage.CacheWriteTokens > 0 {
usage.CacheCreationTokens = ev.Usage.CacheWriteTokens
}
}
}
}

View File

@@ -66,24 +66,6 @@ func TestAccumulateBedrockStream_Converse(t *testing.T) {
require.Equal(t, "pong", completion, "converse text deltas concatenated")
}
// The converse-stream metadata frame's camelCase cache fields must reach the billed cache buckets.
func TestAccumulateBedrockStream_ConverseCacheBuckets(t *testing.T) {
var body bytes.Buffer
body.Write(bedrockFrame(t, "contentBlockDelta", mustJSON(t, map[string]any{"delta": map[string]any{"text": "pong"}})))
body.Write(bedrockFrame(t, "metadata", mustJSON(t, map[string]any{"usage": map[string]any{
"inputTokens": 11, "outputTokens": 3, "totalTokens": 30,
"cacheReadInputTokens": 7, "cacheWriteInputTokens": 9,
}})))
usage, completion := accumulateBedrockStream(body.Bytes())
require.Equal(t, int64(11), usage.InputTokens, "input tokens from metadata frame")
require.Equal(t, int64(3), usage.OutputTokens, "output tokens from metadata frame")
require.Equal(t, int64(7), usage.CachedInputTokens, "cache-read tokens from metadata frame")
require.Equal(t, int64(9), usage.CacheCreationTokens, "cache-write tokens from metadata frame")
require.Equal(t, int64(30), usage.TotalTokens, "provider-reported total wins")
require.Equal(t, "pong", completion)
}
func TestAccumulateBedrockStream_Truncated(t *testing.T) {
// A body cut mid-frame must not panic; partial usage is returned.
full := bedrockFrame(t, "metadata", mustJSON(t, map[string]any{"usage": map[string]any{"inputTokens": 11, "outputTokens": 3}}))

View File

@@ -75,19 +75,8 @@ const (
KeyLLMAttributionGroupID = "llm.attribution_group_id"
KeyLLMAttributionWindowS = "llm.attribution_window_seconds"
// Cost metering (emitted by cost_meter). The four per-bucket keys are the
// base of the breakdown — one per token bucket the provider bills
// separately — and the two aggregates below are derived from them:
// usd_total is their sum, usd_cache is cached_input + cache_creation.
KeyCostUSDInput = "cost.usd_input"
// KeyCostUSDCachedInput is the cost of the cache-read bucket (Anthropic cache_read; OpenAI's discounted cached subset of input).
KeyCostUSDCachedInput = "cost.usd_cached_input"
// KeyCostUSDCacheCreation is the cost of the cache-write bucket. Zero for providers without one.
KeyCostUSDCacheCreation = "cost.usd_cache_creation"
KeyCostUSDOutput = "cost.usd_output"
KeyCostUSDTotal = "cost.usd_total"
// KeyCostUSDCache is the portion of cost.usd_total billed for prompt-cache buckets (cache read/creation, or OpenAI's cached input subset).
KeyCostUSDCache = "cost.usd_cache"
// Cost metering (emitted by cost_meter).
KeyCostUSDTotal = "cost.usd_total"
KeyCostSkipped = "cost.skipped"
// Framework-emitted error markers. Use the mw.<id>.* prefix to

View File

@@ -1,337 +0,0 @@
// Package modeldiscovery fetches and normalizes model catalogs from
// Agent Network provider endpoints. Discovery always runs on the proxy so
// it observes the same network path as inference traffic.
package modeldiscovery
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"path"
"sort"
"strings"
"time"
"unicode"
"unicode/utf8"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/netbird/proxy/internal/roundtrip"
)
const (
SourceOpenAIV1Models = "openai_v1_models"
SourceOllamaAPITags = "ollama_api_tags"
defaultTimeout = 5 * time.Second
maxResponseBytes = 1 << 20 // 1 MiB, after HTTP decompression.
maxUpstreamURLBytes = 4096
maxHeaderNameBytes = 256
maxHeaderValueBytes = 64 << 10
maxModels = 500
maxModelIDBytes = 512
)
// Request contains the provider-owned values management resolved from the
// persisted provider record. Callers must not populate these fields from a
// dashboard-supplied URL or credential.
type Request struct {
UpstreamURL string
AuthHeaderName string
AuthHeaderValue string
SkipTLSVerify bool
AllowOllamaFallback bool
}
// Model is the deliberately small response surface returned to management.
// Arbitrary fields supplied by an upstream never cross the control channel.
type Model struct {
ID string
Label string
}
// Result is a normalized model catalog and the endpoint shape that supplied
// it.
type Result struct {
Models []Model
Source string
}
// Discoverer owns the HTTP client used for provider probes.
type Discoverer struct {
client *http.Client
timeout time.Duration
bodyLimit int64
}
// New constructs a direct-upstream discoverer. The transport is the same
// host-network transport family used by Agent Network inference routes.
func New(logger *log.Logger) *Discoverer {
return newWithTransport(roundtrip.NewDirectOnly(logger))
}
func newWithTransport(transport http.RoundTripper) *Discoverer {
return &Discoverer{
client: &http.Client{
Transport: transport,
// Redirects could move a credentialed request away from the
// persisted provider origin. Discovery never follows them.
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
return http.ErrUseLastResponse
},
},
timeout: defaultTimeout,
bodyLimit: maxResponseBytes,
}
}
// Discover queries the OpenAI-compatible model-list endpoint. Ollama's native
// tags endpoint is attempted only when explicitly enabled and the primary
// endpoint reports that the route does not exist.
func (d *Discoverer) Discover(ctx context.Context, in Request) (Result, error) {
if d == nil || d.client == nil {
return Result{}, errors.New("model discovery client is unavailable")
}
if err := validateRequest(in); err != nil {
return Result{}, err
}
timeout := d.timeout
if timeout <= 0 || timeout > defaultTimeout {
timeout = defaultTimeout
}
probeCtx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
models, err := d.fetchOpenAIModels(probeCtx, in)
if err == nil {
return Result{Models: models, Source: SourceOpenAIV1Models}, nil
}
if !in.AllowOllamaFallback || !isMissingEndpoint(err) {
return Result{}, err
}
models, err = d.fetchOllamaTags(probeCtx, in)
if err != nil {
return Result{}, err
}
return Result{Models: models, Source: SourceOllamaAPITags}, nil
}
func validateRequest(in Request) error {
rawURL := strings.TrimSpace(in.UpstreamURL)
if rawURL == "" {
return errors.New("model discovery upstream URL is required")
}
if len(rawURL) > maxUpstreamURLBytes {
return errors.New("model discovery upstream URL is too long")
}
if len(in.AuthHeaderName) > maxHeaderNameBytes || len(in.AuthHeaderValue) > maxHeaderValueBytes {
return errors.New("model discovery authentication header is too large")
}
if (in.AuthHeaderName == "") != (in.AuthHeaderValue == "") {
return errors.New("model discovery authentication header is incomplete")
}
if strings.ContainsAny(in.AuthHeaderName, "\r\n") || strings.ContainsAny(in.AuthHeaderValue, "\r\n") {
return errors.New("model discovery authentication header is invalid")
}
if in.AuthHeaderName != "" && !strings.EqualFold(in.AuthHeaderName, "Authorization") {
return errors.New("model discovery authentication header is unsupported")
}
return nil
}
func (d *Discoverer) fetchOpenAIModels(ctx context.Context, in Request) ([]Model, error) {
body, err := d.fetch(ctx, in, "v1/models")
if err != nil {
return nil, err
}
var payload struct {
Data *[]struct {
ID string `json:"id"`
} `json:"data"`
}
if err := json.Unmarshal(body, &payload); err != nil || payload.Data == nil {
return nil, errors.New("upstream returned invalid OpenAI model-list JSON")
}
ids := make([]string, 0, len(*payload.Data))
for _, model := range *payload.Data {
ids = append(ids, model.ID)
}
return normalize(ids)
}
func (d *Discoverer) fetchOllamaTags(ctx context.Context, in Request) ([]Model, error) {
body, err := d.fetch(ctx, in, "api/tags")
if err != nil {
return nil, err
}
var payload struct {
Models *[]struct {
Name string `json:"name"`
Model string `json:"model"`
} `json:"models"`
}
if err := json.Unmarshal(body, &payload); err != nil || payload.Models == nil {
return nil, errors.New("upstream returned invalid Ollama tags JSON")
}
ids := make([]string, 0, len(*payload.Models))
for _, model := range *payload.Models {
id := model.Model
if strings.TrimSpace(id) == "" {
id = model.Name
}
ids = append(ids, id)
}
return normalize(ids)
}
func (d *Discoverer) fetch(ctx context.Context, in Request, endpointPath string) ([]byte, error) {
endpoint, err := buildEndpointURL(in.UpstreamURL, endpointPath)
if err != nil {
return nil, err
}
reqCtx := roundtrip.WithDirectUpstream(ctx)
if in.SkipTLSVerify {
reqCtx = roundtrip.WithSkipTLSVerify(reqCtx)
}
req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, endpoint.String(), nil)
if err != nil {
return nil, errors.New("create model discovery request")
}
req.Header.Set("Accept", "application/json")
if in.AuthHeaderName != "" {
// Phase 3 is Ollama-only. Canonicalizing the sole catalog-owned
// credential header keeps the control message from becoming a generic
// arbitrary-header primitive.
req.Header.Set("Authorization", in.AuthHeaderValue)
}
resp, err := d.client.Do(req)
if err != nil {
if errors.Is(ctx.Err(), context.DeadlineExceeded) {
return nil, errors.New("model discovery timed out")
}
if errors.Is(ctx.Err(), context.Canceled) {
return nil, context.Canceled
}
// Deliberately omit the underlying error: net/http errors include the
// internal URL, which should not be reflected through the public API.
return nil, errors.New("model discovery request failed")
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return nil, &upstreamStatusError{statusCode: resp.StatusCode}
}
limit := d.bodyLimit
if limit <= 0 || limit > maxResponseBytes {
limit = maxResponseBytes
}
if resp.ContentLength > limit {
return nil, errors.New("model discovery response is too large")
}
body, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
if err != nil {
return nil, errors.New("read model discovery response")
}
if int64(len(body)) > limit {
return nil, errors.New("model discovery response is too large")
}
return body, nil
}
func buildEndpointURL(rawURL, endpointPath string) (*url.URL, error) {
parsed, err := url.Parse(strings.TrimSpace(rawURL))
if err != nil || parsed.Host == "" || parsed.Hostname() == "" || parsed.Opaque != "" {
return nil, errors.New("model discovery upstream URL is invalid")
}
switch strings.ToLower(parsed.Scheme) {
case "http":
parsed.Scheme = "http"
case "https":
parsed.Scheme = "https"
default:
return nil, errors.New("model discovery upstream URL must use http or https")
}
if parsed.User != nil {
return nil, errors.New("model discovery upstream URL must not contain credentials")
}
// Match Agent Network routing semantics: the static discovery path is
// appended to any persisted base path. Queries and fragments on a provider
// URL are not forwarded to inference and are likewise excluded here.
parsed.Path = "/" + strings.TrimPrefix(path.Join(parsed.Path, endpointPath), "/")
parsed.RawPath = ""
parsed.RawQuery = ""
parsed.ForceQuery = false
parsed.Fragment = ""
return parsed, nil
}
func normalize(ids []string) ([]Model, error) {
if len(ids) > maxModels {
return nil, errors.New("upstream returned too many models")
}
seen := make(map[string]struct{}, len(ids))
normalized := make([]string, 0, len(ids))
for _, raw := range ids {
id := strings.TrimSpace(raw)
if !validModelID(id) {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
normalized = append(normalized, id)
}
sort.Strings(normalized)
models := make([]Model, 0, len(normalized))
for _, id := range normalized {
models = append(models, Model{ID: id, Label: id})
}
return models, nil
}
func validModelID(id string) bool {
if id == "" || len(id) > maxModelIDBytes || !utf8.ValidString(id) {
return false
}
for _, r := range id {
if unicode.IsControl(r) {
return false
}
}
return true
}
type upstreamStatusError struct {
statusCode int
}
func (e *upstreamStatusError) Error() string {
return fmt.Sprintf("upstream returned HTTP %d", e.statusCode)
}
func isMissingEndpoint(err error) bool {
var statusErr *upstreamStatusError
if !errors.As(err, &statusErr) {
return false
}
return statusErr.statusCode == http.StatusNotFound || statusErr.statusCode == http.StatusMethodNotAllowed
}

View File

@@ -1,315 +0,0 @@
package modeldiscovery
import (
"context"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestDiscoverOpenAIModels(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, http.MethodGet, r.Method)
assert.Equal(t, "/gateway/v1/models", r.URL.Path)
assert.Empty(t, r.URL.RawQuery)
assert.Equal(t, "application/json", r.Header.Get("Accept"))
assert.Equal(t, "Bearer secret", r.Header.Get("Authorization"))
_, _ = io.WriteString(w, `{
"object": "list",
"data": [
{"id": " qwen2.5:latest ", "owned_by": "ignored"},
{"id": "llama3.2:latest"},
{"id": "llama3.2:latest"},
{"id": ""},
{"id": "bad\u0000id"}
]
}`)
}))
defer server.Close()
discoverer := newWithTransport(http.DefaultTransport)
result, err := discoverer.Discover(context.Background(), Request{
UpstreamURL: server.URL + "/gateway/?ignored=true#fragment",
AuthHeaderName: "authorization",
AuthHeaderValue: "Bearer secret",
})
require.NoError(t, err)
assert.Equal(t, SourceOpenAIV1Models, result.Source)
assert.Equal(t, []Model{
{ID: "llama3.2:latest", Label: "llama3.2:latest"},
{ID: "qwen2.5:latest", Label: "qwen2.5:latest"},
}, result.Models)
}
func TestDiscoverOllamaFallback(t *testing.T) {
t.Parallel()
var primaryCalls atomic.Int32
var fallbackCalls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/v1/models":
primaryCalls.Add(1)
http.NotFound(w, r)
case "/api/tags":
fallbackCalls.Add(1)
_, _ = io.WriteString(w, `{
"models": [
{"name": "ignored-name", "model": "gemma3:latest"},
{"name": "llama3.2:latest", "model": ""}
]
}`)
default:
http.Error(w, "unexpected path", http.StatusInternalServerError)
}
}))
defer server.Close()
result, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
UpstreamURL: server.URL,
AllowOllamaFallback: true,
})
require.NoError(t, err)
assert.Equal(t, int32(1), primaryCalls.Load())
assert.Equal(t, int32(1), fallbackCalls.Load())
assert.Equal(t, SourceOllamaAPITags, result.Source)
assert.Equal(t, []Model{
{ID: "gemma3:latest", Label: "gemma3:latest"},
{ID: "llama3.2:latest", Label: "llama3.2:latest"},
}, result.Models)
}
func TestDiscoverDoesNotFallbackOnAuthenticationFailure(t *testing.T) {
t.Parallel()
var fallbackCalls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/api/tags" {
fallbackCalls.Add(1)
}
http.Error(w, "secret upstream body", http.StatusUnauthorized)
}))
defer server.Close()
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
UpstreamURL: server.URL,
AllowOllamaFallback: true,
})
require.EqualError(t, err, "upstream returned HTTP 401")
assert.Zero(t, fallbackCalls.Load())
assert.NotContains(t, err.Error(), "secret upstream body")
assert.NotContains(t, err.Error(), server.URL)
}
func TestDiscoverDoesNotFollowRedirectsOrLeakAuth(t *testing.T) {
t.Parallel()
var redirectedCalls atomic.Int32
redirected := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
redirectedCalls.Add(1)
assert.Empty(t, r.Header.Get("Authorization"))
_, _ = io.WriteString(w, `{"data":[]}`)
}))
defer redirected.Close()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, redirected.URL+"/captured", http.StatusFound)
}))
defer server.Close()
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
UpstreamURL: server.URL,
AuthHeaderName: "Authorization",
AuthHeaderValue: "Bearer must-not-leak",
})
require.EqualError(t, err, "upstream returned HTTP 302")
assert.Zero(t, redirectedCalls.Load())
}
func TestDiscoverEnforcesResponseLimit(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, `{"data":[{"id":"llama3.2:latest"}]}`)
}))
defer server.Close()
discoverer := newWithTransport(http.DefaultTransport)
discoverer.bodyLimit = 16
_, err := discoverer.Discover(context.Background(), Request{UpstreamURL: server.URL})
require.EqualError(t, err, "model discovery response is too large")
}
func TestDiscoverEnforcesTimeout(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
select {
case <-r.Context().Done():
case <-time.After(time.Second):
_, _ = io.WriteString(w, `{"data":[]}`)
}
}))
defer server.Close()
discoverer := newWithTransport(http.DefaultTransport)
discoverer.timeout = 30 * time.Millisecond
_, err := discoverer.Discover(context.Background(), Request{UpstreamURL: server.URL})
require.EqualError(t, err, "model discovery timed out")
}
func TestDiscoverHonorsSkipTLSVerification(t *testing.T) {
t.Parallel()
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, `{"data":[]}`)
}))
defer server.Close()
discoverer := New(nil)
_, err := discoverer.Discover(context.Background(), Request{UpstreamURL: server.URL})
require.EqualError(t, err, "model discovery request failed")
result, err := discoverer.Discover(context.Background(), Request{
UpstreamURL: server.URL,
SkipTLSVerify: true,
})
require.NoError(t, err)
assert.Empty(t, result.Models)
}
func TestDiscoverRejectsUnsafeRequestValues(t *testing.T) {
t.Parallel()
tests := []struct {
name string
request Request
wantErr string
}{
{
name: "unsupported scheme",
request: Request{UpstreamURL: "file:///etc/passwd"},
wantErr: "model discovery upstream URL is invalid",
},
{
name: "URL credentials",
request: Request{UpstreamURL: "http://user:pass@example.com"},
wantErr: "model discovery upstream URL must not contain credentials",
},
{
name: "missing hostname",
request: Request{UpstreamURL: "http://:11434"},
wantErr: "model discovery upstream URL is invalid",
},
{
name: "incomplete auth header",
request: Request{
UpstreamURL: "http://example.com",
AuthHeaderName: "Authorization",
},
wantErr: "model discovery authentication header is incomplete",
},
{
name: "header injection",
request: Request{
UpstreamURL: "http://example.com",
AuthHeaderName: "Authorization",
AuthHeaderValue: "Bearer safe\r\nX-Evil: yes",
},
wantErr: "model discovery authentication header is invalid",
},
{
name: "unsupported auth header",
request: Request{
UpstreamURL: "http://example.com",
AuthHeaderName: "X-API-Key",
AuthHeaderValue: "secret",
},
wantErr: "model discovery authentication header is unsupported",
},
{
name: "oversized URL",
request: Request{
UpstreamURL: "http://example.com/" + strings.Repeat("a", maxUpstreamURLBytes),
},
wantErr: "model discovery upstream URL is too long",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
t.Parallel()
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), test.request)
require.EqualError(t, err, test.wantErr)
})
}
}
func TestDiscoverRejectsInvalidPayloads(t *testing.T) {
t.Parallel()
tests := []struct {
name string
status int
body string
wantErr string
}{
{
name: "malformed JSON",
status: http.StatusOK,
body: `{"data":`,
wantErr: "upstream returned invalid OpenAI model-list JSON",
},
{
name: "missing data",
status: http.StatusOK,
body: `{}`,
wantErr: "upstream returned invalid OpenAI model-list JSON",
},
{
name: "too many models",
status: http.StatusOK,
body: modelListJSON(maxModels + 1),
wantErr: "upstream returned too many models",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(test.status)
_, _ = io.WriteString(w, test.body)
}))
defer server.Close()
_, err := newWithTransport(http.DefaultTransport).Discover(context.Background(), Request{
UpstreamURL: server.URL,
})
require.EqualError(t, err, test.wantErr)
})
}
}
func modelListJSON(count int) string {
var body strings.Builder
body.WriteString(`{"data":[`)
for i := range count {
if i > 0 {
body.WriteByte(',')
}
_, _ = fmt.Fprintf(&body, `{"id":"model-%d"}`, i)
}
body.WriteString(`]}`)
return body.String()
}

View File

@@ -271,10 +271,6 @@ func (c *testProxyController) SendServiceUpdateToCluster(_ context.Context, _ st
// noop
}
func (c *testProxyController) DiscoverModels(_ context.Context, _, _ string, _ *proto.ModelDiscoveryRequest) (*proto.ModelDiscoveryResult, error) {
return nil, nbproxy.ErrModelDiscoveryUnavailable
}
func (c *testProxyController) GetOIDCValidationConfig() nbproxy.OIDCValidationConfig {
return nbproxy.OIDCValidationConfig{}
}

View File

@@ -1,311 +0,0 @@
package proxy
import (
"context"
"errors"
"fmt"
"io"
"sync/atomic"
"testing"
"time"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"github.com/netbirdio/netbird/proxy/internal/crowdsec"
"github.com/netbirdio/netbird/proxy/internal/modeldiscovery"
"github.com/netbirdio/netbird/shared/management/proto"
)
type stubModelDiscoverer struct {
started chan modeldiscovery.Request
release <-chan struct{}
result modeldiscovery.Result
err error
}
func (s *stubModelDiscoverer) Discover(ctx context.Context, request modeldiscovery.Request) (modeldiscovery.Result, error) {
if s.started != nil {
select {
case s.started <- request:
case <-ctx.Done():
return modeldiscovery.Result{}, ctx.Err()
}
}
if s.release != nil {
select {
case <-s.release:
case <-ctx.Done():
return modeldiscovery.Result{}, ctx.Err()
}
}
return s.result, s.err
}
type modelDiscoverySyncStream struct {
grpc.ClientStream
ctx context.Context
recv chan *proto.SyncMappingsResponse
sent chan *proto.SyncMappingsRequest
sendWait time.Duration
sending atomic.Int32
overlap atomic.Bool
}
func newModelDiscoverySyncStream(ctx context.Context) *modelDiscoverySyncStream {
return &modelDiscoverySyncStream{
ctx: ctx,
recv: make(chan *proto.SyncMappingsResponse, 16),
sent: make(chan *proto.SyncMappingsRequest, 16),
}
}
func (s *modelDiscoverySyncStream) Send(message *proto.SyncMappingsRequest) error {
if s.sending.Add(1) != 1 {
s.overlap.Store(true)
}
defer s.sending.Add(-1)
if s.sendWait > 0 {
time.Sleep(s.sendWait)
}
select {
case s.sent <- message:
return nil
case <-s.ctx.Done():
return s.ctx.Err()
}
}
func (s *modelDiscoverySyncStream) Recv() (*proto.SyncMappingsResponse, error) {
select {
case message, ok := <-s.recv:
if !ok {
return nil, io.EOF
}
return message, nil
case <-s.ctx.Done():
return nil, s.ctx.Err()
}
}
func (s *modelDiscoverySyncStream) Context() context.Context {
return s.ctx
}
func TestProxyCapabilitiesAdvertiseModelDiscovery(t *testing.T) {
t.Parallel()
server := &Server{
crowdsecRegistry: crowdsec.NewRegistry("", "", log.New().WithField("test", true)),
}
capabilities := server.proxyCapabilities()
require.NotNil(t, capabilities.SupportsModelDiscovery)
assert.True(t, capabilities.GetSupportsModelDiscovery())
}
func TestExecuteModelDiscoveryMapsControlShape(t *testing.T) {
t.Parallel()
discoverer := &stubModelDiscoverer{
result: modeldiscovery.Result{
Source: modeldiscovery.SourceOpenAIV1Models,
Models: []modeldiscovery.Model{
{ID: "llama3.2:latest", Label: "Llama 3.2"},
},
},
}
request := &proto.ModelDiscoveryRequest{
RequestId: "request-1",
UpstreamUrl: "http://ollama.internal:11434",
AuthHeaderName: "Authorization",
AuthHeaderValue: "Bearer secret",
SkipTlsVerify: true,
OllamaFallback: true,
}
result := executeModelDiscovery(context.Background(), discoverer, request)
assert.Equal(t, "request-1", result.GetRequestId())
assert.Equal(t, modeldiscovery.SourceOpenAIV1Models, result.GetSource())
require.Len(t, result.GetModels(), 1)
assert.Equal(t, "llama3.2:latest", result.GetModels()[0].GetId())
assert.Equal(t, "Llama 3.2", result.GetModels()[0].GetLabel())
discoverer.started = make(chan modeldiscovery.Request, 1)
_ = executeModelDiscovery(context.Background(), discoverer, request)
received := <-discoverer.started
assert.Equal(t, request.GetUpstreamUrl(), received.UpstreamURL)
assert.Equal(t, request.GetAuthHeaderName(), received.AuthHeaderName)
assert.Equal(t, request.GetAuthHeaderValue(), received.AuthHeaderValue)
assert.True(t, received.SkipTLSVerify)
assert.True(t, received.AllowOllamaFallback)
}
func TestExecuteModelDiscoveryReturnsSanitizedError(t *testing.T) {
t.Parallel()
result := executeModelDiscovery(context.Background(), &stubModelDiscoverer{
err: errors.New("model discovery request failed"),
}, &proto.ModelDiscoveryRequest{RequestId: "request-error"})
assert.Equal(t, "request-error", result.GetRequestId())
assert.Equal(t, "model discovery request failed", result.GetError())
assert.Empty(t, result.GetModels())
assert.Empty(t, result.GetSource())
}
func TestHandleSyncMappingsStreamRunsDiscoveryOutOfBand(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
release := make(chan struct{})
started := make(chan modeldiscovery.Request, 1)
server := &Server{
Logger: log.New(),
routerReady: closedChan(),
modelDiscoverer: &stubModelDiscoverer{
started: started,
release: release,
result: modeldiscovery.Result{
Source: modeldiscovery.SourceOpenAIV1Models,
Models: []modeldiscovery.Model{{ID: "model-a", Label: "model-a"}},
},
},
}
stream := newModelDiscoverySyncStream(ctx)
stream.sendWait = 10 * time.Millisecond
done := make(chan error, 1)
initialSyncDone := true
go func() {
done <- server.handleSyncMappingsStream(ctx, stream, &initialSyncDone, time.Time{})
}()
stream.recv <- &proto.SyncMappingsResponse{
ModelDiscoveryRequest: &proto.ModelDiscoveryRequest{
RequestId: "request-1",
UpstreamUrl: "http://ollama.internal:11434",
},
}
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("model discovery did not start")
}
// A normal mapping batch must still be acknowledged while the HTTP probe
// is in flight.
stream.recv <- &proto.SyncMappingsResponse{}
select {
case sent := <-stream.sent:
assert.NotNil(t, sent.GetAck())
assert.Nil(t, sent.GetModelDiscoveryResult())
case <-time.After(time.Second):
t.Fatal("mapping ack was blocked by model discovery")
}
close(release)
select {
case sent := <-stream.sent:
result := sent.GetModelDiscoveryResult()
require.NotNil(t, result)
assert.Equal(t, "request-1", result.GetRequestId())
assert.Equal(t, modeldiscovery.SourceOpenAIV1Models, result.GetSource())
assert.Nil(t, sent.GetAck())
case <-time.After(time.Second):
t.Fatal("model discovery result was not sent")
}
assert.False(t, stream.overlap.Load(), "acks and discovery results must use one serialized sender")
close(stream.recv)
require.NoError(t, <-done)
}
func TestHandleSyncMappingsStreamBoundsConcurrentDiscovery(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
release := make(chan struct{})
started := make(chan modeldiscovery.Request, 8)
server := &Server{
Logger: log.New(),
routerReady: closedChan(),
modelDiscoverer: &stubModelDiscoverer{
started: started,
release: release,
},
}
stream := newModelDiscoverySyncStream(ctx)
done := make(chan error, 1)
initialSyncDone := true
go func() {
done <- server.handleSyncMappingsStream(ctx, stream, &initialSyncDone, time.Time{})
}()
for i := range 5 {
stream.recv <- &proto.SyncMappingsResponse{
ModelDiscoveryRequest: &proto.ModelDiscoveryRequest{
RequestId: fmt.Sprintf("request-%d", i),
UpstreamUrl: "http://ollama.internal:11434",
},
}
}
for range 4 {
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("expected four concurrent model discoveries")
}
}
select {
case sent := <-stream.sent:
result := sent.GetModelDiscoveryResult()
require.NotNil(t, result)
assert.Equal(t, "model discovery is busy", result.GetError())
case <-time.After(time.Second):
t.Fatal("fifth discovery did not fail fast")
}
close(release)
for range 4 {
select {
case sent := <-stream.sent:
require.NotNil(t, sent.GetModelDiscoveryResult())
case <-time.After(time.Second):
t.Fatal("in-flight model discovery did not complete")
}
}
close(stream.recv)
require.NoError(t, <-done)
}
func TestHandleSyncMappingsStreamRejectsMixedDiscoveryMessage(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
server := &Server{
Logger: log.New(),
routerReady: closedChan(),
modelDiscoverer: &stubModelDiscoverer{},
}
stream := newModelDiscoverySyncStream(ctx)
stream.recv <- &proto.SyncMappingsResponse{
Mapping: []*proto.ProxyMapping{{Id: "mapping-1"}},
ModelDiscoveryRequest: &proto.ModelDiscoveryRequest{
RequestId: "request-1",
},
}
close(stream.recv)
initialSyncDone := true
err := server.handleSyncMappingsStream(ctx, stream, &initialSyncDone, time.Time{})
require.EqualError(t, err, "model discovery message must not include mapping data")
}

View File

@@ -57,7 +57,6 @@ import (
proxymetrics "github.com/netbirdio/netbird/proxy/internal/metrics"
"github.com/netbirdio/netbird/proxy/internal/middleware"
mwbuiltin "github.com/netbirdio/netbird/proxy/internal/middleware/builtin"
"github.com/netbirdio/netbird/proxy/internal/modeldiscovery"
"github.com/netbirdio/netbird/proxy/internal/netutil"
"github.com/netbirdio/netbird/proxy/internal/proxy"
"github.com/netbirdio/netbird/proxy/internal/restrict"
@@ -79,10 +78,6 @@ type portRouter struct {
cancel context.CancelFunc
}
type providerModelDiscoverer interface {
Discover(context.Context, modeldiscovery.Request) (modeldiscovery.Result, error)
}
type Server struct {
ctx context.Context
mgmtClient proto.ProxyServiceClient
@@ -105,20 +100,16 @@ type Server struct {
// middlewareRegistry is the source of registered middleware factories.
// Concrete middlewares register themselves through init().
middlewareRegistry *middleware.Registry
// modelDiscoverer executes explicit provider model-list probes on the
// proxy's host network. Lazily constructed for normal servers; injectable
// in focused control-stream tests.
modelDiscoverer providerModelDiscoverer
mainRouter *nbtcp.Router
mainPort uint16
udpMu sync.Mutex
udpRelays map[types.ServiceID]*udprelay.Relay
udpRelayWg sync.WaitGroup
portMu sync.RWMutex
portRouters map[uint16]*portRouter
svcPorts map[types.ServiceID][]uint16
lastMappings map[types.ServiceID]*proto.ProxyMapping
portRouterWg sync.WaitGroup
mainRouter *nbtcp.Router
mainPort uint16
udpMu sync.Mutex
udpRelays map[types.ServiceID]*udprelay.Relay
udpRelayWg sync.WaitGroup
portMu sync.RWMutex
portRouters map[uint16]*portRouter
svcPorts map[types.ServiceID][]uint16
lastMappings map[types.ServiceID]*proto.ProxyMapping
portRouterWg sync.WaitGroup
// hijackTracker tracks hijacked connections (e.g. WebSocket upgrades)
// so they can be closed during graceful shutdown, since http.Server.Shutdown
@@ -1286,16 +1277,12 @@ func (s *Server) proxyCapabilities() *proto.ProxyCapabilities {
privateCapability := s.Private
// Always true: this build enforces ProxyMapping.private via the auth middleware.
supportsPrivateService := true
// Model discovery is handled only on SyncMappings. Management also gates
// dispatch on the live connection using this capability.
supportsModelDiscovery := true
return &proto.ProxyCapabilities{
SupportsCustomPorts: &s.SupportsCustomPorts,
RequireSubdomain: &s.RequireSubdomain,
SupportsCrowdsec: &supportsCrowdSec,
Private: &privateCapability,
SupportsPrivateService: &supportsPrivateService,
SupportsModelDiscovery: &supportsModelDiscovery,
}
}
@@ -1362,8 +1349,7 @@ func isSyncUnimplemented(err error) bool {
// handleSyncMappingsStream consumes batches from a bidirectional SyncMappings
// stream, sending an ack after each batch is fully processed. Management waits
// for the ack before sending the next batch, providing application-level
// back-pressure. Model discovery commands are out-of-band: they run with
// bounded concurrency and return a correlated result instead of an ack.
// back-pressure.
func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.ProxyService_SyncMappingsClient, initialSyncDone *bool, connectTime time.Time) error {
select {
case <-s.routerReady:
@@ -1372,29 +1358,6 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
}
tracker := s.newSnapshotTracker(initialSyncDone, connectTime)
discoverer := s.modelDiscoverer
if discoverer == nil {
discoverer = modeldiscovery.New(s.Logger)
}
const maxConcurrentDiscoveries = 4
discoverySlots := make(chan struct{}, maxConcurrentDiscoveries)
discoveryCtx, cancelDiscoveries := context.WithCancel(ctx)
var discoveryWG sync.WaitGroup
var sendMu sync.Mutex
defer func() {
cancelDiscoveries()
discoveryWG.Wait()
}()
// gRPC permits one concurrent sender and one concurrent receiver, but not
// multiple senders. Mapping acks and asynchronous discovery results share
// this serialized send path.
send := func(message *proto.SyncMappingsRequest) error {
sendMu.Lock()
defer sendMu.Unlock()
return stream.Send(message)
}
for {
select {
@@ -1409,46 +1372,6 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
return fmt.Errorf("receive msg: %w", err)
}
if discovery := msg.GetModelDiscoveryRequest(); discovery != nil {
if len(msg.GetMapping()) != 0 || msg.GetInitialSyncComplete() {
return errors.New("model discovery message must not include mapping data")
}
select {
case discoverySlots <- struct{}{}:
discoveryWG.Add(1)
go func() {
defer discoveryWG.Done()
defer func() { <-discoverySlots }()
result := executeModelDiscovery(discoveryCtx, discoverer, discovery)
if discoveryCtx.Err() != nil {
return
}
if err := send(&proto.SyncMappingsRequest{
Msg: &proto.SyncMappingsRequest_ModelDiscoveryResult{
ModelDiscoveryResult: result,
},
}); err != nil {
s.Logger.WithError(err).Debug("failed to send model discovery result")
}
}()
default:
result := &proto.ModelDiscoveryResult{
RequestId: discovery.GetRequestId(),
Error: "model discovery is busy",
}
if err := send(&proto.SyncMappingsRequest{
Msg: &proto.SyncMappingsRequest_ModelDiscoveryResult{
ModelDiscoveryResult: result,
},
}); err != nil {
return fmt.Errorf("send model discovery busy result: %w", err)
}
}
continue
}
batchStart := time.Now()
s.Logger.Debug("Received mapping update, starting processing")
if err := s.processMappingsGuarded(ctx, msg.GetMapping()); err != nil {
@@ -1457,7 +1380,7 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
s.Logger.Debug("Processing mapping update completed")
tracker.recordBatch(ctx, s, msg.GetMapping(), msg.GetInitialSyncComplete(), batchStart)
if err := send(&proto.SyncMappingsRequest{
if err := stream.Send(&proto.SyncMappingsRequest{
Msg: &proto.SyncMappingsRequest_Ack{Ack: &proto.SyncMappingsAck{}},
}); err != nil {
return fmt.Errorf("send ack: %w", err)
@@ -1466,31 +1389,6 @@ func (s *Server) handleSyncMappingsStream(ctx context.Context, stream proto.Prox
}
}
func executeModelDiscovery(ctx context.Context, discoverer providerModelDiscoverer, request *proto.ModelDiscoveryRequest) *proto.ModelDiscoveryResult {
result := &proto.ModelDiscoveryResult{RequestId: request.GetRequestId()}
discovered, err := discoverer.Discover(ctx, modeldiscovery.Request{
UpstreamURL: request.GetUpstreamUrl(),
AuthHeaderName: request.GetAuthHeaderName(),
AuthHeaderValue: request.GetAuthHeaderValue(),
SkipTLSVerify: request.GetSkipTlsVerify(),
AllowOllamaFallback: request.GetOllamaFallback(),
})
if err != nil {
result.Error = err.Error()
return result
}
result.Source = discovered.Source
result.Models = make([]*proto.ModelDiscoveryModel, 0, len(discovered.Models))
for _, model := range discovered.Models {
result.Models = append(result.Models, &proto.ModelDiscoveryModel{
Id: model.ID,
Label: model.Label,
})
}
return result
}
// snapshotTracker accumulates service IDs during the initial snapshot and
// finalises sync state when the complete flag arrives. Used by both
// handleMappingStream and handleSyncMappingsStream so metric emission and

View File

@@ -5138,9 +5138,6 @@ components:
description: Models exposed through this endpoint, with the operator's per-1k input/output prices. Empty means all catalog models are allowed at catalog prices.
items:
$ref: '#/components/schemas/AgentNetworkProviderModel'
has_api_key:
type: boolean
description: Whether an upstream API key is currently stored. The key itself is never returned.
extra_values:
type: object
description: |
@@ -5189,7 +5186,6 @@ components:
- name
- upstream_url
- models
- has_api_key
- enabled
- skip_tls_verification
- metadata_disabled
@@ -5216,8 +5212,7 @@ components:
example: "eu.proxy.netbird.io"
api_key:
type: string
description: |
Upstream provider API key. Sealed at rest on the management server and never returned in responses. Whether a key is accepted or required is declared by the selected catalog provider's `auth_mode`. On update, omit this field to preserve the existing key or send an empty string to clear an optional key.
description: Upstream provider API key. Sealed at rest on the management server and never returned in responses. Required on create; optional on update (omit to keep the existing key).
example: "sk-..."
models:
type: array
@@ -5280,51 +5275,6 @@ components:
- id
- input_per_1k
- output_per_1k
AgentNetworkDiscoveredModel:
type: object
description: A model reported by the provider's persisted upstream endpoint.
properties:
id:
type: string
description: Exact model identifier accepted by the upstream endpoint.
example: "llama3.2:latest"
label:
type: string
description: Human-friendly label reported by the endpoint. Falls back to the model identifier.
example: "llama3.2:latest"
required:
- id
- label
AgentNetworkModelDiscoveryResponse:
type: object
description: |
Result of an explicit model-discovery probe executed by a connected
proxy in the account's selected Agent Network cluster. Discovered
models are not persisted until the provider is updated.
properties:
models:
type: array
description: Normalized, deduplicated models reported by the upstream.
items:
$ref: '#/components/schemas/AgentNetworkDiscoveredModel'
source:
type: string
description: Upstream API shape used to obtain the result.
enum: [openai_v1_models, ollama_api_tags]
example: "openai_v1_models"
proxy_cluster:
type: string
description: Proxy cluster from which the endpoint was queried.
example: "eu.proxy.netbird.io"
request_id:
type: string
description: Correlation identifier for the management-to-proxy probe.
example: "b7ef7da0-78cc-4583-8c4f-55f2ce72015f"
required:
- models
- source
- proxy_cluster
- request_id
AgentNetworkCatalogModel:
type: object
properties:
@@ -5375,11 +5325,6 @@ components:
type: string
description: Default upstream host suggested when adding a provider of this type.
example: "api.openai.com"
auth_mode:
type: string
description: Whether this provider requires, optionally accepts, or does not support an upstream API key.
enum: [required, optional, none]
example: "required"
auth_header_template:
type: string
description: Template the proxy uses to inject the API key (the literal string ${API_KEY} is replaced at request time).
@@ -5401,10 +5346,6 @@ components:
"custom" — generic OpenAI-compatible self-hosted endpoint catch-all.
enum: [provider, gateway, custom]
example: "provider"
supports_model_discovery:
type: boolean
description: Whether models can be loaded from a persisted endpoint through the selected proxy cluster.
example: false
extra_headers:
type: array
description: |
@@ -5423,12 +5364,10 @@ components:
- name
- description
- default_host
- auth_mode
- auth_header_template
- default_content_type
- brand_color
- kind
- supports_model_discovery
- models
AgentNetworkCatalogIdentityInjection:
type: object
@@ -5883,48 +5822,13 @@ components:
total_tokens:
type: integer
format: int64
description: Total tokens consumed, including prompt-cache tokens.
description: Total tokens consumed.
example: 1840
cached_input_tokens:
type: integer
format: int64
description: Input tokens read from the provider's prompt cache. Additive to input_tokens for Anthropic-shape providers; a subset of input_tokens for OpenAI.
example: 0
cache_creation_tokens:
type: integer
format: int64
description: Input tokens written to the provider's prompt cache. Zero for providers without a cache-write bucket.
example: 30528
cost_usd:
type: number
format: double
description: Estimated USD cost of the request.
example: 0.0231
input_cost_usd:
type: number
format: double
description: Cost of the non-cached input tokens. Base component of cost_usd.
example: 0.0048
cached_input_cost_usd:
type: number
format: double
description: Cost of the prompt-cache read tokens. Base component of cost_usd, and part of cache_cost_usd.
example: 0.0015
cache_creation_cost_usd:
type: number
format: double
description: Cost of the prompt-cache write tokens. Base component of cost_usd, and part of cache_cost_usd.
example: 0.1130
output_cost_usd:
type: number
format: double
description: Cost of the output tokens. Base component of cost_usd.
example: 0.0038
cache_cost_usd:
type: number
format: double
description: Portion of cost_usd billed for prompt-cache usage.
example: 0.1145
stream:
type: boolean
description: Whether the request was a streaming completion.
@@ -5948,14 +5852,7 @@ components:
- input_tokens
- output_tokens
- total_tokens
- cached_input_tokens
- cache_creation_tokens
- input_cost_usd
- cached_input_cost_usd
- cache_creation_cost_usd
- output_cost_usd
- cost_usd
- cache_cost_usd
AgentNetworkAccessLogsResponse:
type: object
properties:
@@ -6029,48 +5926,13 @@ components:
total_tokens:
type: integer
format: int64
description: Total tokens across the session, including prompt-cache tokens.
description: Total tokens across the session.
example: 12880
cached_input_tokens:
type: integer
format: int64
description: Total prompt-cache read tokens across the session.
example: 0
cache_creation_tokens:
type: integer
format: int64
description: Total prompt-cache write tokens across the session.
example: 30528
cost_usd:
type: number
format: double
description: Total estimated USD cost across the session.
example: 0.1617
input_cost_usd:
type: number
format: double
description: Total cost of non-cached input tokens across the session.
example: 0.0210
cached_input_cost_usd:
type: number
format: double
description: Total cost of prompt-cache read tokens across the session.
example: 0.0015
cache_creation_cost_usd:
type: number
format: double
description: Total cost of prompt-cache write tokens across the session.
example: 0.1130
output_cost_usd:
type: number
format: double
description: Total cost of output tokens across the session.
example: 0.0262
cache_cost_usd:
type: number
format: double
description: Portion of cost_usd billed for prompt-cache usage across the session.
example: 0.1145
providers:
type: array
items:
@@ -6097,14 +5959,7 @@ components:
- input_tokens
- output_tokens
- total_tokens
- cached_input_tokens
- cache_creation_tokens
- input_cost_usd
- cached_input_cost_usd
- cache_creation_cost_usd
- output_cost_usd
- cost_usd
- cache_cost_usd
- decision
- entries
AgentNetworkAccessLogSessionsResponse:
@@ -6158,61 +6013,19 @@ components:
total_tokens:
type: integer
format: int64
description: Total tokens in the bucket, including prompt-cache tokens.
description: Total tokens in the bucket.
example: 184000
cached_input_tokens:
type: integer
format: int64
description: Total prompt-cache read tokens in the bucket.
example: 20000
cache_creation_tokens:
type: integer
format: int64
description: Total prompt-cache write tokens in the bucket.
example: 45000
input_cost_usd:
type: number
format: double
description: Total cost of non-cached input tokens in the bucket.
example: 1.12
cached_input_cost_usd:
type: number
format: double
description: Total cost of prompt-cache read tokens in the bucket.
example: 0.06
cache_creation_cost_usd:
type: number
format: double
description: Total cost of prompt-cache write tokens in the bucket.
example: 0.36
output_cost_usd:
type: number
format: double
description: Total cost of output tokens in the bucket.
example: 0.77
cost_usd:
type: number
format: double
description: Total estimated USD spend in the bucket.
example: 2.31
cache_cost_usd:
type: number
format: double
description: Portion of cost_usd billed for prompt-cache usage in the bucket.
example: 0.42
required:
- period_start
- input_tokens
- output_tokens
- total_tokens
- cached_input_tokens
- cache_creation_tokens
- input_cost_usd
- cached_input_cost_usd
- cache_creation_cost_usd
- output_cost_usd
- cost_usd
- cache_cost_usd
AgentNetworkSettings:
type: object
description: Per-account Agent Network gateway settings. One row per account; cluster and subdomain are auto-assigned on first provider create and immutable thereafter.
@@ -14089,48 +13902,6 @@ paths:
"$ref": "#/components/responses/not_found"
'500':
"$ref": "#/components/responses/internal_error"
/api/agent-network/providers/{providerId}/discover-models:
post:
summary: Discover models from an Agent Network provider
description: |
Executes a bounded, read-only model-list probe from a connected proxy
in the account's selected Agent Network cluster. The probe uses the
provider's persisted endpoint, TLS setting, and stored credential.
Returned models are not persisted by this operation.
tags: [ Agent Network ]
security:
- BearerAuth: [ ]
- TokenAuth: [ ]
parameters:
- in: path
name: providerId
required: true
schema:
type: string
description: The unique identifier of a persisted Agent Network provider
responses:
'200':
description: Models discovered successfully
content:
application/json:
schema:
$ref: '#/components/schemas/AgentNetworkModelDiscoveryResponse'
'400':
"$ref": "#/components/responses/bad_request"
'401':
"$ref": "#/components/responses/requires_authentication"
'403':
"$ref": "#/components/responses/forbidden"
'404':
"$ref": "#/components/responses/not_found"
'412':
description: The provider cannot be queried or no capable proxy is connected
content: { }
'429':
description: A discovery probe is already running or was requested too recently
content: { }
'500':
"$ref": "#/components/responses/internal_error"
/api/agent-network/policies:
get:
summary: List all Agent Network Policies

View File

@@ -38,27 +38,6 @@ func (e AccessRestrictionsCrowdsecMode) Valid() bool {
}
}
// Defines values for AgentNetworkCatalogProviderAuthMode.
const (
AgentNetworkCatalogProviderAuthModeNone AgentNetworkCatalogProviderAuthMode = "none"
AgentNetworkCatalogProviderAuthModeOptional AgentNetworkCatalogProviderAuthMode = "optional"
AgentNetworkCatalogProviderAuthModeRequired AgentNetworkCatalogProviderAuthMode = "required"
)
// Valid indicates whether the value is a known member of the AgentNetworkCatalogProviderAuthMode enum.
func (e AgentNetworkCatalogProviderAuthMode) Valid() bool {
switch e {
case AgentNetworkCatalogProviderAuthModeNone:
return true
case AgentNetworkCatalogProviderAuthModeOptional:
return true
case AgentNetworkCatalogProviderAuthModeRequired:
return true
default:
return false
}
}
// Defines values for AgentNetworkCatalogProviderKind.
const (
AgentNetworkCatalogProviderKindCustom AgentNetworkCatalogProviderKind = "custom"
@@ -98,24 +77,6 @@ func (e AgentNetworkConsumptionDimensionKind) Valid() bool {
}
}
// Defines values for AgentNetworkModelDiscoveryResponseSource.
const (
AgentNetworkModelDiscoveryResponseSourceOllamaApiTags AgentNetworkModelDiscoveryResponseSource = "ollama_api_tags"
AgentNetworkModelDiscoveryResponseSourceOpenaiV1Models AgentNetworkModelDiscoveryResponseSource = "openai_v1_models"
)
// Valid indicates whether the value is a known member of the AgentNetworkModelDiscoveryResponseSource enum.
func (e AgentNetworkModelDiscoveryResponseSource) Valid() bool {
switch e {
case AgentNetworkModelDiscoveryResponseSourceOllamaApiTags:
return true
case AgentNetworkModelDiscoveryResponseSourceOpenaiV1Models:
return true
default:
return false
}
}
// Defines values for CreateAzureIntegrationRequestHost.
const (
CreateAzureIntegrationRequestHostMicrosoftCom CreateAzureIntegrationRequestHost = "microsoft.com"
@@ -1771,21 +1732,6 @@ type AccountSettings struct {
// AgentNetworkAccessLog One per-request agent-network (LLM) access log entry with flattened, queryable LLM dimensions.
type AgentNetworkAccessLog struct {
// CacheCostUsd Portion of cost_usd billed for prompt-cache usage.
CacheCostUsd float64 `json:"cache_cost_usd"`
// CacheCreationCostUsd Cost of the prompt-cache write tokens. Base component of cost_usd, and part of cache_cost_usd.
CacheCreationCostUsd float64 `json:"cache_creation_cost_usd"`
// CacheCreationTokens Input tokens written to the provider's prompt cache. Zero for providers without a cache-write bucket.
CacheCreationTokens int64 `json:"cache_creation_tokens"`
// CachedInputCostUsd Cost of the prompt-cache read tokens. Base component of cost_usd, and part of cache_cost_usd.
CachedInputCostUsd float64 `json:"cached_input_cost_usd"`
// CachedInputTokens Input tokens read from the provider's prompt cache. Additive to input_tokens for Anthropic-shape providers; a subset of input_tokens for OpenAI.
CachedInputTokens int64 `json:"cached_input_tokens"`
// CostUsd Estimated USD cost of the request.
CostUsd float64 `json:"cost_usd"`
@@ -1807,9 +1753,6 @@ type AgentNetworkAccessLog struct {
// Id Unique identifier for the access log entry.
Id string `json:"id"`
// InputCostUsd Cost of the non-cached input tokens. Base component of cost_usd.
InputCostUsd float64 `json:"input_cost_usd"`
// InputTokens Input (prompt) tokens consumed.
InputTokens int64 `json:"input_tokens"`
@@ -1819,9 +1762,6 @@ type AgentNetworkAccessLog struct {
// Model Requested LLM model.
Model *string `json:"model,omitempty"`
// OutputCostUsd Cost of the output tokens. Base component of cost_usd.
OutputCostUsd float64 `json:"output_cost_usd"`
// OutputTokens Output (completion) tokens produced.
OutputTokens int64 `json:"output_tokens"`
@@ -1861,7 +1801,7 @@ type AgentNetworkAccessLog struct {
// Timestamp Timestamp when the request was made.
Timestamp time.Time `json:"timestamp"`
// TotalTokens Total tokens consumed, including prompt-cache tokens.
// TotalTokens Total tokens consumed.
TotalTokens int64 `json:"total_tokens"`
// UserId NetBird user id of the authenticated caller, if applicable.
@@ -1870,21 +1810,6 @@ type AgentNetworkAccessLog struct {
// AgentNetworkAccessLogSession A session-grouped view of agent-network access logs — all requests sharing a session id (or a single session-less request) folded into one summary plus its ordered entries.
type AgentNetworkAccessLogSession struct {
// CacheCostUsd Portion of cost_usd billed for prompt-cache usage across the session.
CacheCostUsd float64 `json:"cache_cost_usd"`
// CacheCreationCostUsd Total cost of prompt-cache write tokens across the session.
CacheCreationCostUsd float64 `json:"cache_creation_cost_usd"`
// CacheCreationTokens Total prompt-cache write tokens across the session.
CacheCreationTokens int64 `json:"cache_creation_tokens"`
// CachedInputCostUsd Total cost of prompt-cache read tokens across the session.
CachedInputCostUsd float64 `json:"cached_input_cost_usd"`
// CachedInputTokens Total prompt-cache read tokens across the session.
CachedInputTokens int64 `json:"cached_input_tokens"`
// CostUsd Total estimated USD cost across the session.
CostUsd float64 `json:"cost_usd"`
@@ -1900,18 +1825,12 @@ type AgentNetworkAccessLogSession struct {
// GroupIds Union of the authorising group ids across the session's entries.
GroupIds *[]string `json:"group_ids,omitempty"`
// InputCostUsd Total cost of non-cached input tokens across the session.
InputCostUsd float64 `json:"input_cost_usd"`
// InputTokens Total input (prompt) tokens across the session.
InputTokens int64 `json:"input_tokens"`
// Models Distinct models seen in the session.
Models *[]string `json:"models,omitempty"`
// OutputCostUsd Total cost of output tokens across the session.
OutputCostUsd float64 `json:"output_cost_usd"`
// OutputTokens Total output (completion) tokens across the session.
OutputTokens int64 `json:"output_tokens"`
@@ -1927,7 +1846,7 @@ type AgentNetworkAccessLogSession struct {
// StartedAt Timestamp of the session's earliest request.
StartedAt time.Time `json:"started_at"`
// TotalTokens Total tokens across the session, including prompt-cache tokens.
// TotalTokens Total tokens across the session.
TotalTokens int64 `json:"total_tokens"`
// UserId NetBird user id of the session's caller.
@@ -2077,9 +1996,6 @@ type AgentNetworkCatalogProvider struct {
// AuthHeaderTemplate Template the proxy uses to inject the API key (the literal string ${API_KEY} is replaced at request time).
AuthHeaderTemplate string `json:"auth_header_template"`
// AuthMode Whether this provider requires, optionally accepts, or does not support an upstream API key.
AuthMode AgentNetworkCatalogProviderAuthMode `json:"auth_mode"`
// BrandColor Hex brand color used to render the provider badge in the dashboard.
BrandColor string `json:"brand_color"`
@@ -2112,14 +2028,8 @@ type AgentNetworkCatalogProvider struct {
// Name Display name for the provider.
Name string `json:"name"`
// SupportsModelDiscovery Whether models can be loaded from a persisted endpoint through the selected proxy cluster.
SupportsModelDiscovery bool `json:"supports_model_discovery"`
}
// AgentNetworkCatalogProviderAuthMode Whether this provider requires, optionally accepts, or does not support an upstream API key.
type AgentNetworkCatalogProviderAuthMode string
// AgentNetworkCatalogProviderKind Presentation grouping for the provider Select on the dashboard.
// "provider" — first-party vendor API (OpenAI, Anthropic, …); the upstream is the model itself.
// "gateway" — routing/aggregation layer in front of multiple providers (LiteLLM, Portkey, …); typically pairs with NetBird identity stamping.
@@ -2156,15 +2066,6 @@ type AgentNetworkConsumption struct {
// AgentNetworkConsumptionDimensionKind Whether this row counts a single end user or a single source group across every member.
type AgentNetworkConsumptionDimensionKind string
// AgentNetworkDiscoveredModel A model reported by the provider's persisted upstream endpoint.
type AgentNetworkDiscoveredModel struct {
// Id Exact model identifier accepted by the upstream endpoint.
Id string `json:"id"`
// Label Human-friendly label reported by the endpoint. Falls back to the model identifier.
Label string `json:"label"`
}
// AgentNetworkGuardrail defines model for AgentNetworkGuardrail.
type AgentNetworkGuardrail struct {
// Checks Guardrail check parameters. Each entry has an `enabled` flag plus per-check configuration; disabled entries are inert.
@@ -2212,26 +2113,6 @@ type AgentNetworkGuardrailRequest struct {
Name string `json:"name"`
}
// AgentNetworkModelDiscoveryResponse Result of an explicit model-discovery probe executed by a connected
// proxy in the account's selected Agent Network cluster. Discovered
// models are not persisted until the provider is updated.
type AgentNetworkModelDiscoveryResponse struct {
// Models Normalized, deduplicated models reported by the upstream.
Models []AgentNetworkDiscoveredModel `json:"models"`
// ProxyCluster Proxy cluster from which the endpoint was queried.
ProxyCluster string `json:"proxy_cluster"`
// RequestId Correlation identifier for the management-to-proxy probe.
RequestId string `json:"request_id"`
// Source Upstream API shape used to obtain the result.
Source AgentNetworkModelDiscoveryResponseSource `json:"source"`
}
// AgentNetworkModelDiscoveryResponseSource Upstream API shape used to obtain the result.
type AgentNetworkModelDiscoveryResponseSource string
// AgentNetworkPolicy defines model for AgentNetworkPolicy.
type AgentNetworkPolicy struct {
// CreatedAt Timestamp when the policy was created.
@@ -2337,9 +2218,6 @@ type AgentNetworkProvider struct {
// ExtraValues Operator-typed values for catalog-declared extra headers. Keys are wire header names (e.g. `x-portkey-config`); values are the strings the proxy stamps on every upstream request to this provider. Catalog (AgentNetworkCatalogProvider.extra_headers) declares which keys are accepted; values not declared by the catalog are ignored at synth time. Empty / missing values mean no header stamped.
ExtraValues *map[string]string `json:"extra_values,omitempty"`
// HasApiKey Whether an upstream API key is currently stored. The key itself is never returned.
HasApiKey bool `json:"has_api_key"`
// Id Provider ID
Id string `json:"id"`
@@ -2385,7 +2263,7 @@ type AgentNetworkProviderModel struct {
// AgentNetworkProviderRequest defines model for AgentNetworkProviderRequest.
type AgentNetworkProviderRequest struct {
// ApiKey Upstream provider API key. Sealed at rest on the management server and never returned in responses. Whether a key is accepted or required is declared by the selected catalog provider's `auth_mode`. On update, omit this field to preserve the existing key or send an empty string to clear an optional key.
// ApiKey Upstream provider API key. Sealed at rest on the management server and never returned in responses. Required on create; optional on update (omit to keep the existing key).
ApiKey *string `json:"api_key,omitempty"`
// BootstrapCluster Proxy cluster used to bootstrap the per-account agent-network endpoint when the first provider is created. Ignored on subsequent creates and on updates because the cluster is pinned on the account-level Settings row.
@@ -2469,40 +2347,19 @@ type AgentNetworkSettingsRequest struct {
// AgentNetworkUsageBucket One aggregated agent-network usage time bucket (UTC). The bucket width is set by the request's granularity.
type AgentNetworkUsageBucket struct {
// CacheCostUsd Portion of cost_usd billed for prompt-cache usage in the bucket.
CacheCostUsd float64 `json:"cache_cost_usd"`
// CacheCreationCostUsd Total cost of prompt-cache write tokens in the bucket.
CacheCreationCostUsd float64 `json:"cache_creation_cost_usd"`
// CacheCreationTokens Total prompt-cache write tokens in the bucket.
CacheCreationTokens int64 `json:"cache_creation_tokens"`
// CachedInputCostUsd Total cost of prompt-cache read tokens in the bucket.
CachedInputCostUsd float64 `json:"cached_input_cost_usd"`
// CachedInputTokens Total prompt-cache read tokens in the bucket.
CachedInputTokens int64 `json:"cached_input_tokens"`
// CostUsd Total estimated USD spend in the bucket.
CostUsd float64 `json:"cost_usd"`
// InputCostUsd Total cost of non-cached input tokens in the bucket.
InputCostUsd float64 `json:"input_cost_usd"`
// InputTokens Total input (prompt) tokens in the bucket.
InputTokens int64 `json:"input_tokens"`
// OutputCostUsd Total cost of output tokens in the bucket.
OutputCostUsd float64 `json:"output_cost_usd"`
// OutputTokens Total output (completion) tokens in the bucket.
OutputTokens int64 `json:"output_tokens"`
// PeriodStart Start of the bucket in YYYY-MM-DD (UTC) — the day, the week start (Monday), or the month start, depending on granularity.
PeriodStart string `json:"period_start"`
// TotalTokens Total tokens in the bucket, including prompt-cache tokens.
// TotalTokens Total tokens in the bucket.
TotalTokens int64 `json:"total_tokens"`
}

View File

@@ -11,6 +11,9 @@ import (
log "github.com/sirupsen/logrus"
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/proto"
@@ -35,17 +38,17 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
Network: decodeAccountNetwork(full.Network),
AccountSettings: decodeAccountSettings(full.AccountSettings),
CustomZoneDomain: full.CustomZoneDomain,
Peers: make(map[string]*types.ComponentPeer, len(full.Peers)),
Groups: make(map[string]*types.ComponentGroup, len(full.Groups)),
Peers: make(map[string]*nbpeer.Peer, len(full.Peers)),
Groups: make(map[string]*types.Group, len(full.Groups)),
Policies: make([]*types.Policy, 0, len(full.Policies)),
Routes: make([]*nbroute.Route, 0, len(full.Routes)),
NameServerGroups: make([]*nbdns.NameServerGroup, 0, len(full.NameserverGroups)),
AllDNSRecords: decodeSimpleRecords(full.AllDnsRecords),
AccountZones: decodeCustomZones(full.AccountZones),
ResourcePoliciesMap: make(map[string][]*types.Policy),
RoutersMap: make(map[string]map[string]*types.ComponentRouter),
NetworkResources: make([]*types.ComponentResource, 0, len(full.NetworkResources)),
RouterPeers: make(map[string]*types.ComponentPeer),
RoutersMap: make(map[string]map[string]*routerTypes.NetworkRouter),
NetworkResources: make([]*resourceTypes.NetworkResource, 0, len(full.NetworkResources)),
RouterPeers: make(map[string]*nbpeer.Peer),
AllowedUserIDs: stringSliceToSet(full.AllowedUserIds),
PostureFailedPeers: make(map[string]map[string]struct{}, len(full.PostureFailedPeers)),
GroupIDToUserIDs: make(map[string][]string, len(full.GroupIdToUserIds)),
@@ -98,7 +101,7 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
log.WithField("peer idx", idx).Error("unrecognized peer idx during decoding")
}
}
group := &types.ComponentGroup{
group := &types.Group{
ID: groupID,
PublicID: gc.Id,
Peers: peerIDs,
@@ -148,7 +151,7 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
// Phase 7: routers_map (outer key = network seq id, inner key = peer-id
// reconstructed from peer_index). Synthesized network id is "net_<seq>".
for networkID, list := range full.RoutersMap {
inner := make(map[string]*types.ComponentRouter, len(list.Entries))
inner := make(map[string]*routerTypes.NetworkRouter, len(list.Entries))
for _, entry := range list.Entries {
if !entry.PeerIndexSet {
continue
@@ -158,7 +161,8 @@ func DecodeEnvelope(env *proto.NetworkMapEnvelope) (*types.NetworkMapComponents,
continue
}
peerID := peerIDByIndex[entry.PeerIndex]
inner[peerID] = &types.ComponentRouter{
inner[peerID] = &routerTypes.NetworkRouter{
ID: "",
NetworkID: networkID,
PublicID: entry.Id,
Peer: peerID,
@@ -260,22 +264,40 @@ func decodeAccountSettings(as *proto.AccountSettingsCompact) *types.AccountSetti
}
}
func decodePeerCompact(pc *proto.PeerCompact, peerID string) *types.ComponentPeer {
peer := &types.ComponentPeer{
func decodePeerCompact(pc *proto.PeerCompact, peerID string) *nbpeer.Peer {
var caps []int32
if pc.SupportsSourcePrefixes {
caps = append(caps, nbpeer.PeerCapabilitySourcePrefixes)
}
if pc.SupportsIpv6 {
caps = append(caps, nbpeer.PeerCapabilityIPv6Overlay)
}
peer := &nbpeer.Peer{
ID: peerID,
Key: peerID,
SSHKey: string(pc.SshPubKey),
SSHEnabled: pc.SshEnabled,
DNSLabel: pc.DnsLabel,
LoginExpirationEnabled: pc.LoginExpirationEnabled,
AgentVersion: pc.AgentVersion,
SupportsSourcePrefixes: pc.SupportsSourcePrefixes,
SupportsIPv6: pc.SupportsIpv6,
ServerSSHAllowed: pc.ServerSshAllowed,
AddedWithSSOLogin: pc.AddedWithSsoLogin,
Meta: nbpeer.PeerSystemMeta{
WtVersion: pc.AgentVersion,
Capabilities: caps,
Flags: nbpeer.Flags{
ServerSSHAllowed: pc.ServerSshAllowed,
},
},
}
if pc.AddedWithSsoLogin {
// Set a non-empty UserID so (*Peer).AddedWithSSOLogin() returns true.
// The original UserID isn't on the wire; the value is intentionally
// visibly synthetic so any future consumer that mistakes UserID for a
// real account user xid won't silently match (or worse, write the
// sentinel into a downstream record).
peer.UserID = "<env-sso>"
}
if pc.LastLoginUnixNano != 0 {
peer.LastLogin = time.Unix(0, pc.LastLoginUnixNano)
t := time.Unix(0, pc.LastLoginUnixNano)
peer.LastLogin = &t
}
switch len(pc.Ip) {
case 4:
@@ -402,14 +424,14 @@ func decodeNameServerGroupRaw(nsg *proto.NameServerGroupRaw) *nbdns.NameServerGr
return out
}
func decodeNetworkResource(nr *proto.NetworkResourceRaw) *types.ComponentResource {
out := &types.ComponentResource{
func decodeNetworkResource(nr *proto.NetworkResourceRaw) *resourceTypes.NetworkResource {
out := &resourceTypes.NetworkResource{
ID: nr.Id,
PublicID: nr.Id,
NetworkID: nr.NetworkSeq,
Name: nr.Name,
Description: nr.Description,
Type: types.ComponentResourceType(nr.Type),
Type: resourceTypes.NetworkResourceType(nr.Type),
Address: nr.Address,
Domain: nr.DomainValue,
Enabled: nr.Enabled,

View File

@@ -20,6 +20,7 @@ import (
nbdns "github.com/netbirdio/netbird/dns"
"net/netip"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
nbroute "github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/management/types"
@@ -273,7 +274,7 @@ func ToProtocolDNSConfig(update nbdns.Config, cache DNSConfigCache, forwardPort
// AppendRemotePeerConfig appends typed peers as proto.RemotePeerConfig
// entries to dst and returns the result.
func AppendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*types.ComponentPeer, dnsName string, includeIPv6 bool) []*proto.RemotePeerConfig {
func AppendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*nbpeer.Peer, dnsName string, includeIPv6 bool) []*proto.RemotePeerConfig {
for _, rPeer := range peers {
allowedIPs := []string{rPeer.IP.String() + "/32"}
if includeIPv6 && rPeer.IPv6.IsValid() {
@@ -284,7 +285,7 @@ func AppendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*types.Compon
AllowedIps: allowedIPs,
SshConfig: &proto.SSHConfig{SshPubKey: []byte(rPeer.SSHKey)},
Fqdn: rPeer.FQDN(dnsName),
AgentVersion: rPeer.AgentVersion,
AgentVersion: rPeer.Meta.WtVersion,
})
}
return dst

View File

@@ -54,8 +54,8 @@ func EnvelopeToNetworkMap(ctx context.Context, env *proto.NetworkMapEnvelope, lo
}
components.PeerID = canonicalKey
includeIPv6 := localPeer.SupportsIPv6 && localPeer.IPv6.IsValid()
useSourcePrefixes := localPeer.SupportsSourcePrefixes
includeIPv6 := localPeer.SupportsIPv6() && localPeer.IPv6.IsValid()
useSourcePrefixes := localPeer.SupportsSourcePrefixes()
typedNM := components.Calculate(ctx)

View File

@@ -13,6 +13,7 @@ import (
goproto "google.golang.org/protobuf/proto"
mgmtgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/types"
nbnetworkmap "github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/proto"
@@ -143,14 +144,14 @@ func TestDecodeEnvelope_MalformedWgKeyPeerSkipped(t *testing.T) {
func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) {
ctx := context.Background()
peers := map[string]*types.ComponentPeer{}
peers := map[string]*nbpeer.Peer{}
for i, id := range []string{"peer-T", "peer-S", "peer-ALL", "peer-O"} {
peers[id] = &types.ComponentPeer{
ID: id,
Key: randomWgKey(t),
IP: netip.AddrFrom4([4]byte{100, 64, 0, byte(i + 1)}),
DNSLabel: id,
AgentVersion: "0.40.0",
peers[id] = &nbpeer.Peer{
ID: id,
Key: randomWgKey(t),
IP: netip.AddrFrom4([4]byte{100, 64, 0, byte(i + 1)}),
DNSLabel: id,
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
}
@@ -164,7 +165,7 @@ func TestEnvelopeRoundTrip_AllGroupShortCircuitParity(t *testing.T) {
AccountSettings: &types.AccountSettingsInfo{},
DNSSettings: &types.DNSSettings{},
Peers: peers,
Groups: map[string]*types.ComponentGroup{
Groups: map[string]*types.Group{
"g-src": {ID: "g-src", PublicID: "1", Name: "staff", Peers: []string{"peer-T", "peer-S"}},
"g-all": {ID: "g-all", PublicID: "2", Name: "All", Peers: []string{"peer-ALL"}},
"g-two": {ID: "g-two", PublicID: "3", Name: "second", Peers: []string{"peer-T", "peer-O"}},
@@ -231,22 +232,22 @@ func buildSmokeComponents(t *testing.T) (*types.NetworkMapComponents, string) {
peerAKey := randomWgKey(t)
peerBKey := randomWgKey(t)
peerA := &types.ComponentPeer{
ID: "peer-A",
Key: peerAKey,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peerA",
AgentVersion: "0.40.0",
peerA := &nbpeer.Peer{
ID: "peer-A",
Key: peerAKey,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 1}),
DNSLabel: "peerA",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
peerB := &types.ComponentPeer{
ID: "peer-B",
Key: peerBKey,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
DNSLabel: "peerB",
AgentVersion: "0.40.0",
peerB := &nbpeer.Peer{
ID: "peer-B",
Key: peerBKey,
IP: netip.AddrFrom4([4]byte{100, 64, 0, 2}),
DNSLabel: "peerB",
Meta: nbpeer.PeerSystemMeta{WtVersion: "0.40.0"},
}
group := &types.ComponentGroup{
group := &types.Group{
ID: "group-all", PublicID: "1", Name: "All",
Peers: []string{"peer-A", "peer-B"},
}
@@ -273,11 +274,11 @@ func buildSmokeComponents(t *testing.T) (*types.NetworkMapComponents, string) {
},
AccountSettings: &types.AccountSettingsInfo{},
DNSSettings: &types.DNSSettings{},
Peers: map[string]*types.ComponentPeer{
Peers: map[string]*nbpeer.Peer{
"peer-A": peerA,
"peer-B": peerB,
},
Groups: map[string]*types.ComponentGroup{
Groups: map[string]*types.Group{
"group-all": group,
},
Policies: []*types.Policy{policy},

File diff suppressed because it is too large Load Diff

View File

@@ -73,10 +73,6 @@ message ProxyCapabilities {
optional bool private = 4;
// Whether the proxy enforces ProxyMapping.private (fails closed on ValidateTunnelPeer failure). Management MUST NOT stream private mappings to proxies that don't claim this.
optional bool supports_private_service = 5;
// Whether the proxy can execute correlated Agent Network model-discovery
// requests delivered over SyncMappings. Management must not send discovery
// requests to proxies that omit or disable this capability.
optional bool supports_model_discovery = 6;
}
// GetMappingUpdateRequest is sent to initialise a mapping stream.
@@ -424,8 +420,6 @@ message SyncMappingsRequest {
oneof msg {
SyncMappingsInit init = 1;
SyncMappingsAck ack = 2;
// Correlated response to a ModelDiscoveryRequest received from management.
ModelDiscoveryResult model_discovery_result = 3;
}
}
@@ -449,9 +443,6 @@ message SyncMappingsResponse {
repeated ProxyMapping mapping = 1;
// initial_sync_complete is set on the last message of the initial snapshot.
bool initial_sync_complete = 2;
// An out-of-band model-discovery operation. This is sent only after the
// initial snapshot and only to proxies advertising supports_model_discovery.
ModelDiscoveryRequest model_discovery_request = 3;
}
// CheckLLMPolicyLimitsRequest carries the resolved caller identity and the
@@ -510,33 +501,3 @@ message RecordLLMUsageRequest {
message RecordLLMUsageResponse {
}
// ModelDiscoveryRequest asks a proxy to list models from a persisted provider
// endpoint using the proxy host's network path. request_id is assigned by
// management and echoed verbatim in ModelDiscoveryResult.
message ModelDiscoveryRequest {
string request_id = 1;
string upstream_url = 2;
string auth_header_name = 3;
string auth_header_value = 4;
bool skip_tls_verify = 5;
// When true, the proxy may fall back from the OpenAI-compatible
// GET /v1/models endpoint to Ollama's native GET /api/tags endpoint.
bool ollama_fallback = 6;
}
// ModelDiscoveryModel is the normalized model shape returned to management.
// Arbitrary upstream fields never cross the proxy control channel.
message ModelDiscoveryModel {
string id = 1;
string label = 2;
}
// ModelDiscoveryResult completes one ModelDiscoveryRequest. error is empty on
// success and contains only a sanitized diagnostic on failure.
message ModelDiscoveryResult {
string request_id = 1;
repeated ModelDiscoveryModel models = 2;
// Stable source label, currently openai_v1_models or ollama_api_tags.
string source = 3;
string error = 4;
}

View File

@@ -1,103 +0,0 @@
package types
import (
"net/netip"
"time"
)
// ComponentPeer is the self-contained peer representation used by
// NetworkMapComponents and the calculated NetworkMap. It carries exactly the
// subset of peer data that crosses the components wire format, so the shared
// calculation layer stays independent of the management server's domain
// types.
type ComponentPeer struct {
ID string
Key string
IP netip.Addr
IPv6 netip.Addr
DNSLabel string
SSHKey string
SSHEnabled bool
ServerSSHAllowed bool
AgentVersion string
SupportsSourcePrefixes bool
SupportsIPv6 bool
LoginExpirationEnabled bool
AddedWithSSOLogin bool
LastLogin time.Time
}
// FQDN returns the peer's FQDN combined of the peer's DNS label and the system's DNS domain.
func (p *ComponentPeer) FQDN(dnsDomain string) string {
if dnsDomain == "" {
return ""
}
return p.DNSLabel + "." + dnsDomain
}
// LoginExpired indicates whether the peer's login has expired, mirroring the
// server-side peer semantics: only SSO-added peers with login expiration
// enabled can expire.
func (p *ComponentPeer) LoginExpired(expiresIn time.Duration) (bool, time.Duration) {
if !p.AddedWithSSOLogin || !p.LoginExpirationEnabled {
return false, 0
}
timeLeft := time.Until(p.LastLogin.Add(expiresIn))
return timeLeft <= 0, timeLeft
}
// GroupAllName is the reserved name of the default group that contains every peer in an account.
const GroupAllName = "All"
// ComponentGroup is the self-contained group representation used by
// NetworkMapComponents: just the membership view the network-map calculation
// needs, without the server's storage fields.
type ComponentGroup struct {
ID string
PublicID string
Name string
Peers []string
}
// IsGroupAll checks if the group is a default "All" group.
func (g *ComponentGroup) IsGroupAll() bool {
return g.Name == GroupAllName
}
// ComponentRouter is the self-contained network-router representation used by
// NetworkMapComponents.
type ComponentRouter struct {
NetworkID string
PublicID string
Peer string
PeerGroups []string
Masquerade bool
Metric int
Enabled bool
}
// ComponentResourceType mirrors the network-resource type enum on the
// components wire format.
type ComponentResourceType string
const (
ComponentResourceHost ComponentResourceType = "host"
ComponentResourceSubnet ComponentResourceType = "subnet"
ComponentResourceDomain ComponentResourceType = "domain"
)
// ComponentResource is the self-contained network-resource representation
// used by NetworkMapComponents.
type ComponentResource struct {
ID string
PublicID string
NetworkID string
AccountID string
Name string
Description string
Type ComponentResourceType
Address string
Domain string
Prefix netip.Prefix
Enabled bool
}

View File

@@ -3,6 +3,8 @@ package types
import (
"strconv"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/version"
)
@@ -46,8 +48,8 @@ func portsIncludesSSH(ports []string) bool {
}
// ExpandPortsAndRanges expands Ports and PortRanges of a rule into individual firewall rules.
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *ComponentPeer) []*FirewallRule {
features := peerSupportedFirewallFeatures(peer.AgentVersion)
func ExpandPortsAndRanges(base FirewallRule, rule *PolicyRule, peer *nbpeer.Peer) []*FirewallRule {
features := peerSupportedFirewallFeatures(peer.Meta.WtVersion)
var expanded []*FirewallRule
@@ -104,8 +106,8 @@ func isPortInRule(portString string, portInt uint16, rule *FirewallRule) bool {
return rule.Port == portString || (rule.PortRange.Start <= portInt && portInt <= rule.PortRange.End)
}
func shouldCheckRulesForNativeSSH(supportsNative bool, rule *PolicyRule, peer *ComponentPeer) bool {
return supportsNative && peer.SSHEnabled && peer.ServerSSHAllowed && rule.Protocol == PolicyRuleProtocolTCP
func shouldCheckRulesForNativeSSH(supportsNative bool, rule *PolicyRule, peer *nbpeer.Peer) bool {
return supportsNative && peer.SSHEnabled && peer.Meta.Flags.ServerSSHAllowed && rule.Protocol == PolicyRuleProtocolTCP
}
func peerSupportedFirewallFeatures(peerVer string) supportedFeatures {
@@ -115,13 +117,13 @@ func peerSupportedFirewallFeatures(peerVer string) supportedFeatures {
var features supportedFeatures
meetMinVer, err := version.MeetsMinVersion(firewallRuleMinNativeSSHVer, peerVer)
meetMinVer, err := posture.MeetsMinVersion(firewallRuleMinNativeSSHVer, peerVer)
features.nativeSSH = err == nil && meetMinVer
if features.nativeSSH {
features.portRanges = true
} else {
meetMinVer, err = version.MeetsMinVersion(firewallRuleMinPortRangesVer, peerVer)
meetMinVer, err = posture.MeetsMinVersion(firewallRuleMinPortRangesVer, peerVer)
features.portRanges = err == nil && meetMinVer
}

View File

@@ -9,6 +9,7 @@ import (
log "github.com/sirupsen/logrus"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
nbroute "github.com/netbirdio/netbird/route"
)
@@ -50,7 +51,7 @@ func (r *FirewallRule) Equal(other *FirewallRule) bool {
// For static routes, source ranges match the destination family (v4 or v6).
// For dynamic routes (domain-based), separate v4 and v6 rules are generated
// so the routing peer's forwarding chain allows both address families.
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*ComponentPeer, direction int, includeIPv6 bool) []*RouteFirewallRule {
func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int, includeIPv6 bool) []*RouteFirewallRule {
rulesExists := make(map[string]struct{})
rules := make([]*RouteFirewallRule, 0)
@@ -106,7 +107,7 @@ func GenerateRouteFirewallRules(ctx context.Context, route *nbroute.Route, rule
}
// splitPeerSourcesByFamily separates peer IPs into v4 (/32) and v6 (/128) source ranges.
func splitPeerSourcesByFamily(groupPeers []*ComponentPeer) (v4, v6 []string) {
func splitPeerSourcesByFamily(groupPeers []*nbpeer.Peer) (v4, v6 []string) {
v4 = make([]string, 0, len(groupPeers))
v6 = make([]string, 0, len(groupPeers))
for _, peer := range groupPeers {

View File

@@ -8,12 +8,13 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/domain"
)
func TestSplitPeerSourcesByFamily(t *testing.T) {
peers := []*ComponentPeer{
peers := []*nbpeer.Peer{
{
IP: netip.MustParseAddr("100.64.0.1"),
IPv6: netip.MustParseAddr("fd00::1"),
@@ -35,7 +36,7 @@ func TestSplitPeerSourcesByFamily(t *testing.T) {
}
func TestGenerateRouteFirewallRules_V4Route(t *testing.T) {
peers := []*ComponentPeer{
peers := []*nbpeer.Peer{
{
IP: netip.MustParseAddr("100.64.0.1"),
IPv6: netip.MustParseAddr("fd00::1"),
@@ -64,7 +65,7 @@ func TestGenerateRouteFirewallRules_V4Route(t *testing.T) {
}
func TestGenerateRouteFirewallRules_V6Route(t *testing.T) {
peers := []*ComponentPeer{
peers := []*nbpeer.Peer{
{
IP: netip.MustParseAddr("100.64.0.1"),
IPv6: netip.MustParseAddr("fd00::1"),
@@ -92,7 +93,7 @@ func TestGenerateRouteFirewallRules_V6Route(t *testing.T) {
}
func TestGenerateRouteFirewallRules_DynamicRoute_DualStack(t *testing.T) {
peers := []*ComponentPeer{
peers := []*nbpeer.Peer{
{
IP: netip.MustParseAddr("100.64.0.1"),
IPv6: netip.MustParseAddr("fd00::1"),
@@ -125,7 +126,7 @@ func TestGenerateRouteFirewallRules_DynamicRoute_DualStack(t *testing.T) {
}
func TestGenerateRouteFirewallRules_DynamicRoute_NoV6Peers(t *testing.T) {
peers := []*ComponentPeer{
peers := []*nbpeer.Peer{
{IP: netip.MustParseAddr("100.64.0.1")},
{IP: netip.MustParseAddr("100.64.0.2")},
}
@@ -149,7 +150,7 @@ func TestGenerateRouteFirewallRules_DynamicRoute_NoV6Peers(t *testing.T) {
}
func TestGenerateRouteFirewallRules_IncludeIPv6False(t *testing.T) {
peers := []*ComponentPeer{
peers := []*nbpeer.Peer{
{
IP: netip.MustParseAddr("100.64.0.1"),
IPv6: netip.MustParseAddr("fd00::1"),

View File

@@ -2,6 +2,7 @@ package types
import (
"github.com/netbirdio/netbird/management/server/integration_reference"
"github.com/netbirdio/netbird/management/server/networks/resources/types"
)
const (
@@ -67,6 +68,10 @@ func (g *Group) EventMeta() map[string]any {
return map[string]any{"name": g.Name}
}
func (g *Group) EventMetaResource(resource *types.NetworkResource) map[string]any {
return map[string]any{"name": g.Name, "id": g.ID, "resource_name": resource.Name, "resource_id": resource.ID, "resource_type": resource.Type}
}
func (g *Group) Copy() *Group {
group := &Group{
ID: g.ID,
@@ -90,39 +95,14 @@ func (g *Group) HasPeers() bool {
return len(g.Peers) > 0
}
// GroupAllName is the reserved name of the default group that contains every peer in an account.
const GroupAllName = "All"
// IsGroupAll checks if the group is a default "All" group.
func (g *Group) IsGroupAll() bool {
return g.Name == GroupAllName
}
// ToComponent converts the group to its self-contained components
// representation. The Peers slice is shared, not copied — components are
// treated as immutable snapshots. Returns nil for a nil group.
func (g *Group) ToComponent() *ComponentGroup {
if g == nil {
return nil
}
return &ComponentGroup{
ID: g.ID,
PublicID: g.PublicID,
Name: g.Name,
Peers: g.Peers,
}
}
// GroupsToComponent converts an id-keyed group map to its components
// representation, preserving nil entries.
func GroupsToComponent(groups map[string]*Group) map[string]*ComponentGroup {
if groups == nil {
return nil
}
out := make(map[string]*ComponentGroup, len(groups))
for id, g := range groups {
out[id] = g.ToComponent()
}
return out
}
// AddPeer adds peerID to Peers if not present, returning true if added.
func (g *Group) AddPeer(peerID string) bool {
if peerID == "" {

View File

@@ -15,6 +15,8 @@ import (
"golang.org/x/exp/maps"
nbdns "github.com/netbirdio/netbird/dns"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/util"
"github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/netbirdio/netbird/shared/management/status"
@@ -37,11 +39,11 @@ const (
)
type NetworkMap struct {
Peers []*ComponentPeer
Peers []*nbpeer.Peer
Network *Network
Routes []*route.Route
DNSConfig nbdns.Config
OfflinePeers []*ComponentPeer
OfflinePeers []*nbpeer.Peer
FirewallRules []*FirewallRule
RoutesFirewallRules []*RouteFirewallRule
ForwardingRules []*ForwardingRule
@@ -51,46 +53,15 @@ type NetworkMap struct {
func (nm *NetworkMap) Merge(other *NetworkMap) {
nm.Peers = mergeUniquePeersByID(nm.Peers, other.Peers)
nm.Routes = mergeUnique(nm.Routes, other.Routes)
nm.Routes = util.MergeUnique(nm.Routes, other.Routes)
nm.OfflinePeers = mergeUniquePeersByID(nm.OfflinePeers, other.OfflinePeers)
nm.FirewallRules = mergeUnique(nm.FirewallRules, other.FirewallRules)
nm.RoutesFirewallRules = mergeUnique(nm.RoutesFirewallRules, other.RoutesFirewallRules)
nm.ForwardingRules = mergeUnique(nm.ForwardingRules, other.ForwardingRules)
nm.FirewallRules = util.MergeUnique(nm.FirewallRules, other.FirewallRules)
nm.RoutesFirewallRules = util.MergeUnique(nm.RoutesFirewallRules, other.RoutesFirewallRules)
nm.ForwardingRules = util.MergeUnique(nm.ForwardingRules, other.ForwardingRules)
}
type comparableObject[T any] interface {
Equal(other T) bool
}
func mergeUnique[T comparableObject[T]](arr1, arr2 []T) []T {
var result []T
for _, item := range arr1 {
if !containsEqual(result, item) {
result = append(result, item)
}
}
for _, item := range arr2 {
if !containsEqual(result, item) {
result = append(result, item)
}
}
return result
}
func containsEqual[T comparableObject[T]](slice []T, element T) bool {
for _, item := range slice {
if item.Equal(element) {
return true
}
}
return false
}
func mergeUniquePeersByID(peers1, peers2 []*ComponentPeer) []*ComponentPeer {
result := make(map[string]*ComponentPeer)
func mergeUniquePeersByID(peers1, peers2 []*nbpeer.Peer) []*nbpeer.Peer {
result := make(map[string]*nbpeer.Peer)
for _, peer := range peers1 {
result[peer.ID] = peer
}

View File

@@ -12,6 +12,9 @@ import (
"github.com/netbirdio/netbird/client/ssh/auth"
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/route"
"github.com/netbirdio/netbird/shared/management/domain"
)
@@ -24,22 +27,22 @@ type NetworkMapComponents struct {
DNSSettings *DNSSettings
CustomZoneDomain string
Peers map[string]*ComponentPeer
Groups map[string]*ComponentGroup
Peers map[string]*nbpeer.Peer
Groups map[string]*Group
Policies []*Policy
Routes []*route.Route
NameServerGroups []*nbdns.NameServerGroup
AllDNSRecords []nbdns.SimpleRecord
AccountZones []nbdns.CustomZone
ResourcePoliciesMap map[string][]*Policy
RoutersMap map[string]map[string]*ComponentRouter
NetworkResources []*ComponentResource
RoutersMap map[string]map[string]*routerTypes.NetworkRouter
NetworkResources []*resourceTypes.NetworkResource
GroupIDToUserIDs map[string][]string
AllowedUserIDs map[string]struct{}
PostureFailedPeers map[string]map[string]struct{}
RouterPeers map[string]*ComponentPeer
RouterPeers map[string]*nbpeer.Peer
// NetworkXIDToPublicID maps Network.ID (xid) → PublicID.
// Consumed by the envelope encoder to
@@ -75,15 +78,15 @@ func EmptyNetworkMapComponents(nm *NetworkMapComponents) *NetworkMapComponents {
return nm
}
func (c *NetworkMapComponents) GetPeerInfo(peerID string) *ComponentPeer {
func (c *NetworkMapComponents) GetPeerInfo(peerID string) *nbpeer.Peer {
return c.Peers[peerID]
}
func (c *NetworkMapComponents) GetRouterPeerInfo(peerID string) *ComponentPeer {
func (c *NetworkMapComponents) GetRouterPeerInfo(peerID string) *nbpeer.Peer {
return c.RouterPeers[peerID]
}
func (c *NetworkMapComponents) GetGroupInfo(groupID string) *ComponentGroup {
func (c *NetworkMapComponents) GetGroupInfo(groupID string) *Group {
return c.Groups[groupID]
}
@@ -139,7 +142,7 @@ func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap {
includeIPv6 := false
if p := c.Peers[targetPeerID]; p != nil {
includeIPv6 = p.SupportsIPv6 && p.IPv6.IsValid()
includeIPv6 = p.SupportsIPv6() && p.IPv6.IsValid()
}
routesUpdate := filterAndExpandRoutes(c.getRoutesToSync(targetPeerID, peersToConnect, peerGroups), includeIPv6)
routesFirewallRules := c.getPeerRoutesFirewallRules(ctx, targetPeerID, includeIPv6)
@@ -197,7 +200,7 @@ func (c *NetworkMapComponents) IsEmpty() bool {
return c.empty
}
func (c *NetworkMapComponents) getPeerConnectionResources(targetPeerID string) ([]*ComponentPeer, []*FirewallRule, map[string]map[string]struct{}, bool) {
func (c *NetworkMapComponents) getPeerConnectionResources(targetPeerID string) ([]*nbpeer.Peer, []*FirewallRule, map[string]map[string]struct{}, bool) {
targetPeer := c.GetPeerInfo(targetPeerID)
if targetPeer == nil {
return nil, nil, nil, false
@@ -217,7 +220,7 @@ func (c *NetworkMapComponents) getPeerConnectionResources(targetPeerID string) (
continue
}
var sourcePeers, destinationPeers []*ComponentPeer
var sourcePeers, destinationPeers []*nbpeer.Peer
var peerInSources, peerInDestinations bool
if rule.SourceResource.Type == ResourceTypePeer && rule.SourceResource.ID != "" {
@@ -300,13 +303,13 @@ func (c *NetworkMapComponents) getAllowedUserIDs() map[string]struct{} {
return make(map[string]struct{})
}
func (c *NetworkMapComponents) connResourcesGenerator(targetPeer *ComponentPeer) (func(*PolicyRule, []*ComponentPeer, int), func() ([]*ComponentPeer, []*FirewallRule)) {
func (c *NetworkMapComponents) connResourcesGenerator(targetPeer *nbpeer.Peer) (func(*PolicyRule, []*nbpeer.Peer, int), func() ([]*nbpeer.Peer, []*FirewallRule)) {
rulesExists := make(map[string]struct{})
peersExists := make(map[string]struct{})
rules := make([]*FirewallRule, 0)
peers := make([]*ComponentPeer, 0)
peers := make([]*nbpeer.Peer, 0)
return func(rule *PolicyRule, groupPeers []*ComponentPeer, direction int) {
return func(rule *PolicyRule, groupPeers []*nbpeer.Peer, direction int) {
protocol := rule.Protocol
if protocol == PolicyRuleProtocolNetbirdSSH {
protocol = PolicyRuleProtocolTCP
@@ -358,15 +361,15 @@ func (c *NetworkMapComponents) connResourcesGenerator(targetPeer *ComponentPeer)
PortsJoined: portsJoined,
})
}
}, func() ([]*ComponentPeer, []*FirewallRule) {
}, func() ([]*nbpeer.Peer, []*FirewallRule) {
return peers, rules
}
}
func (c *NetworkMapComponents) getAllPeersFromGroups(groups []string, peerID string, sourcePostureChecksIDs []string) ([]*ComponentPeer, bool) {
func (c *NetworkMapComponents) getAllPeersFromGroups(groups []string, peerID string, sourcePostureChecksIDs []string) ([]*nbpeer.Peer, bool) {
peerInGroups := false
uniquePeerIDs := c.getUniquePeerIDsFromGroupsIDs(groups)
filteredPeers := make([]*ComponentPeer, 0, len(uniquePeerIDs))
filteredPeers := make([]*nbpeer.Peer, 0, len(uniquePeerIDs))
for _, p := range uniquePeerIDs {
peerInfo := c.GetPeerInfo(p)
@@ -418,22 +421,22 @@ func (c *NetworkMapComponents) getUniquePeerIDsFromGroupsIDs(groups []string) []
return ids
}
func (c *NetworkMapComponents) getPeerFromResource(resource Resource, peerID string) ([]*ComponentPeer, bool) {
func (c *NetworkMapComponents) getPeerFromResource(resource Resource, peerID string) ([]*nbpeer.Peer, bool) {
if resource.ID == peerID {
return []*ComponentPeer{}, true
return []*nbpeer.Peer{}, true
}
peerInfo := c.GetPeerInfo(resource.ID)
if peerInfo == nil {
return []*ComponentPeer{}, false
return []*nbpeer.Peer{}, false
}
return []*ComponentPeer{peerInfo}, false
return []*nbpeer.Peer{peerInfo}, false
}
func (c *NetworkMapComponents) filterPeersByLoginExpiration(aclPeers []*ComponentPeer) ([]*ComponentPeer, []*ComponentPeer) {
peersToConnect := make([]*ComponentPeer, 0, len(aclPeers))
var expiredPeers []*ComponentPeer
func (c *NetworkMapComponents) filterPeersByLoginExpiration(aclPeers []*nbpeer.Peer) ([]*nbpeer.Peer, []*nbpeer.Peer) {
peersToConnect := make([]*nbpeer.Peer, 0, len(aclPeers))
var expiredPeers []*nbpeer.Peer
for _, p := range aclPeers {
expired, _ := p.LoginExpired(c.AccountSettings.PeerLoginExpiration)
@@ -515,7 +518,7 @@ func filterAndExpandRoutes(routes []*route.Route, includeIPv6 bool) []*route.Rou
return filtered
}
func (c *NetworkMapComponents) getRoutesToSync(peerID string, aclPeers []*ComponentPeer, peerGroups LookupMap) []*route.Route {
func (c *NetworkMapComponents) getRoutesToSync(peerID string, aclPeers []*nbpeer.Peer, peerGroups LookupMap) []*route.Route {
routes, peerDisabledRoutes := c.getRoutingPeerRoutes(peerID)
peerRoutesMembership := make(LookupMap)
for _, r := range append(routes, peerDisabledRoutes...) {
@@ -729,7 +732,7 @@ func (c *NetworkMapComponents) getRouteFirewallRules(ctx context.Context, peerID
return fwRules
}
func (c *NetworkMapComponents) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}) []*ComponentPeer {
func (c *NetworkMapComponents) getRulePeers(rule *PolicyRule, postureChecks []string, peerID string, distributionPeers map[string]struct{}) []*nbpeer.Peer {
distPeersWithPolicy := make(map[string]struct{})
for _, id := range rule.Sources {
group := c.GetGroupInfo(id)
@@ -756,7 +759,7 @@ func (c *NetworkMapComponents) getRulePeers(rule *PolicyRule, postureChecks []st
}
}
distributionGroupPeers := make([]*ComponentPeer, 0, len(distPeersWithPolicy))
distributionGroupPeers := make([]*nbpeer.Peer, 0, len(distPeersWithPolicy))
for pID := range distPeersWithPolicy {
peerInfo := c.GetPeerInfo(pID)
if peerInfo == nil {
@@ -796,8 +799,8 @@ func (c *NetworkMapComponents) getNetworkResourcesRoutesToSync(peerID string) (b
func (c *NetworkMapComponents) processResourcePolicies(
peerID string,
resource *ComponentResource,
networkRoutingPeers map[string]*ComponentRouter,
resource *resourceTypes.NetworkResource,
networkRoutingPeers map[string]*routerTypes.NetworkRouter,
addSourcePeers bool,
allSourcePeers map[string]struct{},
) []*route.Route {
@@ -830,7 +833,7 @@ func (c *NetworkMapComponents) getResourcePolicyPeers(policy *Policy) []string {
return c.getUniquePeerIDsFromGroupsIDs(policy.SourceGroups())
}
func (c *NetworkMapComponents) getNetworkResourcesRoutes(resource *ComponentResource, peerID string, router *ComponentRouter) []*route.Route {
func (c *NetworkMapComponents) getNetworkResourcesRoutes(resource *resourceTypes.NetworkResource, peerID string, router *routerTypes.NetworkRouter) []*route.Route {
resourceAppliedPolicies := c.ResourcePoliciesMap[resource.ID]
var routes []*route.Route
@@ -844,7 +847,7 @@ func (c *NetworkMapComponents) getNetworkResourcesRoutes(resource *ComponentReso
return routes
}
func (c *NetworkMapComponents) networkResourceToRoute(resource *ComponentResource, peer *ComponentPeer, router *ComponentRouter) *route.Route {
func (c *NetworkMapComponents) networkResourceToRoute(resource *resourceTypes.NetworkResource, peer *nbpeer.Peer, router *routerTypes.NetworkRouter) *route.Route {
r := &route.Route{
ID: route.ID(resource.ID + ":" + peer.ID),
AccountID: resource.AccountID,
@@ -858,7 +861,7 @@ func (c *NetworkMapComponents) networkResourceToRoute(resource *ComponentResourc
Description: resource.Description,
}
if resource.Type == ComponentResourceHost || resource.Type == ComponentResourceSubnet {
if resource.Type == resourceTypes.Host || resource.Type == resourceTypes.Subnet {
r.Network = resource.Prefix
r.NetworkType = route.IPv4Network
@@ -867,7 +870,7 @@ func (c *NetworkMapComponents) networkResourceToRoute(resource *ComponentResourc
}
}
if resource.Type == ComponentResourceDomain {
if resource.Type == resourceTypes.Domain {
domainList, err := domain.FromStringList([]string{resource.Domain})
if err == nil {
r.Domains = domainList
@@ -945,11 +948,11 @@ func (c *NetworkMapComponents) getPoliciesSourcePeers(policies []*Policy) map[st
func (c *NetworkMapComponents) addNetworksRoutingPeers(
networkResourcesRoutes []*route.Route,
peerID string,
peersToConnect []*ComponentPeer,
expiredPeers []*ComponentPeer,
peersToConnect []*nbpeer.Peer,
expiredPeers []*nbpeer.Peer,
isRouter bool,
sourcePeers map[string]struct{},
) []*ComponentPeer {
) []*nbpeer.Peer {
networkRoutesPeers := make(map[string]struct{}, len(networkResourcesRoutes))
for _, r := range networkResourcesRoutes {
@@ -999,8 +1002,8 @@ type FirewallRuleContext struct {
PortsJoined string
}
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *ComponentPeer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
if !peer.IPv6.IsValid() || !targetPeer.SupportsIPv6 || !targetPeer.IPv6.IsValid() {
func AppendIPv6FirewallRule(rules []*FirewallRule, rulesExists map[string]struct{}, peer, targetPeer *nbpeer.Peer, rule *PolicyRule, rc FirewallRuleContext) []*FirewallRule {
if !peer.IPv6.IsValid() || !targetPeer.SupportsIPv6() || !targetPeer.IPv6.IsValid() {
return rules
}

View File

@@ -2,6 +2,9 @@ package types
import (
nbdns "github.com/netbirdio/netbird/dns"
resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types"
routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/route"
)
@@ -18,7 +21,7 @@ type NetworkMapComponentsCompact struct {
DNSSettings *DNSSettings
CustomZoneDomain string
AllPeers []*ComponentPeer
AllPeers []*nbpeer.Peer
PeerIndexes []int
RouterPeerIndexes []int
@@ -31,8 +34,8 @@ type NetworkMapComponentsCompact struct {
AllDNSRecords []nbdns.SimpleRecord
AccountZones []nbdns.CustomZone
RoutersMap map[string]map[string]*ComponentRouter
NetworkResources []*ComponentResource
RoutersMap map[string]map[string]*routerTypes.NetworkRouter
NetworkResources []*resourceTypes.NetworkResource
GroupIDToUserIDs map[string][]string
AllowedUserIDs map[string]struct{}
@@ -41,7 +44,7 @@ type NetworkMapComponentsCompact struct {
func (c *NetworkMapComponents) ToCompact() *NetworkMapComponentsCompact {
peerToIndex := make(map[string]int)
var allPeers []*ComponentPeer
var allPeers []*nbpeer.Peer
for id, peer := range c.Peers {
if _, exists := peerToIndex[id]; !exists {
@@ -147,7 +150,7 @@ func (c *NetworkMapComponents) ToCompact() *NetworkMapComponentsCompact {
}
func (c *NetworkMapComponentsCompact) ToFull() *NetworkMapComponents {
peers := make(map[string]*ComponentPeer, len(c.PeerIndexes))
peers := make(map[string]*nbpeer.Peer, len(c.PeerIndexes))
for _, idx := range c.PeerIndexes {
if idx >= 0 && idx < len(c.AllPeers) {
peer := c.AllPeers[idx]
@@ -155,7 +158,7 @@ func (c *NetworkMapComponentsCompact) ToFull() *NetworkMapComponents {
}
}
routerPeers := make(map[string]*ComponentPeer, len(c.RouterPeerIndexes))
routerPeers := make(map[string]*nbpeer.Peer, len(c.RouterPeerIndexes))
for _, idx := range c.RouterPeerIndexes {
if idx >= 0 && idx < len(c.AllPeers) {
peer := c.AllPeers[idx]
@@ -163,7 +166,7 @@ func (c *NetworkMapComponentsCompact) ToFull() *NetworkMapComponents {
}
}
groups := make(map[string]*ComponentGroup, len(c.Groups))
groups := make(map[string]*Group, len(c.Groups))
for id, gc := range c.Groups {
peerIDs := make([]string, 0, len(gc.PeerIndexes))
for _, idx := range gc.PeerIndexes {
@@ -171,7 +174,7 @@ func (c *NetworkMapComponentsCompact) ToFull() *NetworkMapComponents {
peerIDs = append(peerIDs, c.AllPeers[idx].ID)
}
}
groups[id] = &ComponentGroup{
groups[id] = &Group{
ID: id,
Name: gc.Name,
Peers: peerIDs,

View File

@@ -24,7 +24,7 @@ import (
)
// ErrSharedSockStopped indicates that shared socket has been stopped
var ErrSharedSockStopped = fmt.Errorf("shared socket stopped")
var ErrSharedSockStopped = fmt.Errorf("shared socked stopped")
// SharedSocket is a net.PacketConn that initiates two raw sockets (ipv4 and ipv6) and listens to UDP packets filtered
// by BPF instructions (e.g., IncomingSTUNFilter that checks and sends only STUN packets to the listeners (ReadFrom)).

View File

@@ -71,30 +71,6 @@ func NetbirdCommit() string {
return revision
}
// sanitizeVersion removes anything after the pre-release tag (e.g., "-dev", "-alpha", etc.)
func sanitizeVersion(version string) string {
parts := strings.Split(version, "-")
return parts[0]
}
// MeetsMinVersion checks if the peer's version meets or exceeds the minimum required version
func MeetsMinVersion(minVer, peerVer string) (bool, error) {
peerVer = sanitizeVersion(peerVer)
minVer = sanitizeVersion(minVer)
peerNBVer, err := v.NewVersion(peerVer)
if err != nil {
return false, err
}
constraints, err := v.NewConstraint(">= " + minVer)
if err != nil {
return false, err
}
return constraints.Check(peerNBVer), nil
}
// IsDevelopmentVersion reports whether the given version string identifies
// a non-release / development build. It is the single source of truth for
// "is this a dev build" checks across the codebase; use it instead of

View File

@@ -1,10 +1,6 @@
package version
import (
"testing"
"github.com/stretchr/testify/assert"
)
import "testing"
func TestIsDevelopmentVersion(t *testing.T) {
tests := []struct {
@@ -30,68 +26,3 @@ func TestIsDevelopmentVersion(t *testing.T) {
})
}
}
func TestMeetsMinVersion(t *testing.T) {
tests := []struct {
name string
minVer string
peerVer string
want bool
wantErr bool
}{
{
name: "Peer version greater than min version",
minVer: "0.26.0",
peerVer: "0.60.1",
want: true,
wantErr: false,
},
{
name: "Peer version equals min version",
minVer: "1.0.0",
peerVer: "1.0.0",
want: true,
wantErr: false,
},
{
name: "Peer version less than min version",
minVer: "1.0.0",
peerVer: "0.9.9",
want: false,
wantErr: false,
},
{
name: "Peer version with pre-release tag greater than min version",
minVer: "1.0.0",
peerVer: "1.0.1-alpha",
want: true,
wantErr: false,
},
{
name: "Invalid peer version format",
minVer: "1.0.0",
peerVer: "dev",
want: false,
wantErr: true,
},
{
name: "Invalid min version format",
minVer: "invalid.version",
peerVer: "1.0.0",
want: false,
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := MeetsMinVersion(tt.minVer, tt.peerVer)
if tt.wantErr {
assert.Error(t, err)
} else {
assert.NoError(t, err)
}
assert.Equal(t, tt.want, got)
})
}
}