mirror of
https://github.com/netbirdio/netbird.git
synced 2026-07-27 02:41:28 +02:00
Compare commits
1 Commits
feat/agent
...
revert/com
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7ed3737cda |
9
.github/workflows/agent-network-e2e.yml
vendored
9
.github/workflows/agent-network-e2e.yml
vendored
@@ -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 }}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(*)",
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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": {}}},
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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{{
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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})
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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}}))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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{}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
126
proxy/server.go
126
proxy/server.go
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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 == "" {
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)).
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user