diff --git a/shared/management/client/rest/agentnetwork.go b/shared/management/client/rest/agentnetwork.go new file mode 100644 index 000000000..70f3691ac --- /dev/null +++ b/shared/management/client/rest/agentnetwork.go @@ -0,0 +1,381 @@ +package rest + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + + "github.com/netbirdio/netbird/shared/management/http/api" +) + +// AgentNetworkAPI APIs for the Agent Network (AI/LLM gateway), do not use directly +// see more: https://docs.netbird.io/api/resources/agent-network +type AgentNetworkAPI struct { + c *Client +} + +// ListCatalogProviders lists the catalog of supported upstream AI providers +// (openai_api, anthropic_api, bedrock_api, ...) with their default models and +// pricing, used to prefill provider create forms. +func (a *AgentNetworkAPI) ListCatalogProviders(ctx context.Context) ([]api.AgentNetworkCatalogProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/catalog/providers", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkCatalogProvider](resp) + return ret, err +} + +// ListProviders lists all Agent Network providers +func (a *AgentNetworkAPI) ListProviders(ctx context.Context) ([]api.AgentNetworkProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/providers", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkProvider](resp) + return ret, err +} + +// GetProvider gets Agent Network provider info +func (a *AgentNetworkAPI) GetProvider(ctx context.Context, providerID string) (*api.AgentNetworkProvider, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/providers/"+providerID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// CreateProvider creates a new Agent Network provider. Set +// request.BootstrapCluster on the account's first provider to bootstrap the +// per-account gateway endpoint (alternatively bootstrap via UpdateSettings +// with a cluster). +func (a *AgentNetworkAPI) CreateProvider(ctx context.Context, request api.PostApiAgentNetworkProvidersJSONRequestBody) (*api.AgentNetworkProvider, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/providers", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// UpdateProvider updates an Agent Network provider. Omitted optional fields +// (api_key, models, extra_values, toggles, identity headers) keep their +// stored values. +func (a *AgentNetworkAPI) UpdateProvider(ctx context.Context, providerID string, request api.PutApiAgentNetworkProvidersProviderIdJSONRequestBody) (*api.AgentNetworkProvider, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/providers/"+providerID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkProvider](resp) + return &ret, err +} + +// DeleteProvider deletes an Agent Network provider. Fails while any policy +// still references the provider — detach it first. +func (a *AgentNetworkAPI) DeleteProvider(ctx context.Context, providerID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/providers/"+providerID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListPolicies lists all Agent Network policies +func (a *AgentNetworkAPI) ListPolicies(ctx context.Context) ([]api.AgentNetworkPolicy, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/policies", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkPolicy](resp) + return ret, err +} + +// GetPolicy gets Agent Network policy info +func (a *AgentNetworkAPI) GetPolicy(ctx context.Context, policyID string) (*api.AgentNetworkPolicy, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/policies/"+policyID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// CreatePolicy creates a new Agent Network policy +func (a *AgentNetworkAPI) CreatePolicy(ctx context.Context, request api.PostApiAgentNetworkPoliciesJSONRequestBody) (*api.AgentNetworkPolicy, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/policies", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// UpdatePolicy updates an Agent Network policy +func (a *AgentNetworkAPI) UpdatePolicy(ctx context.Context, policyID string, request api.PutApiAgentNetworkPoliciesPolicyIdJSONRequestBody) (*api.AgentNetworkPolicy, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/policies/"+policyID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkPolicy](resp) + return &ret, err +} + +// DeletePolicy deletes an Agent Network policy +func (a *AgentNetworkAPI) DeletePolicy(ctx context.Context, policyID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/policies/"+policyID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListGuardrails lists all Agent Network guardrails +func (a *AgentNetworkAPI) ListGuardrails(ctx context.Context) ([]api.AgentNetworkGuardrail, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/guardrails", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkGuardrail](resp) + return ret, err +} + +// GetGuardrail gets Agent Network guardrail info +func (a *AgentNetworkAPI) GetGuardrail(ctx context.Context, guardrailID string) (*api.AgentNetworkGuardrail, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/guardrails/"+guardrailID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// CreateGuardrail creates a new Agent Network guardrail +func (a *AgentNetworkAPI) CreateGuardrail(ctx context.Context, request api.PostApiAgentNetworkGuardrailsJSONRequestBody) (*api.AgentNetworkGuardrail, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/guardrails", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// UpdateGuardrail updates an Agent Network guardrail +func (a *AgentNetworkAPI) UpdateGuardrail(ctx context.Context, guardrailID string, request api.PutApiAgentNetworkGuardrailsGuardrailIdJSONRequestBody) (*api.AgentNetworkGuardrail, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/guardrails/"+guardrailID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkGuardrail](resp) + return &ret, err +} + +// DeleteGuardrail deletes an Agent Network guardrail +func (a *AgentNetworkAPI) DeleteGuardrail(ctx context.Context, guardrailID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/guardrails/"+guardrailID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// ListBudgetRules lists all account-level Agent Network budget rules +func (a *AgentNetworkAPI) ListBudgetRules(ctx context.Context) ([]api.AgentNetworkBudgetRule, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/budget-rules", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[[]api.AgentNetworkBudgetRule](resp) + return ret, err +} + +// GetBudgetRule gets Agent Network budget rule info +func (a *AgentNetworkAPI) GetBudgetRule(ctx context.Context, ruleID string) (*api.AgentNetworkBudgetRule, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/budget-rules/"+ruleID, nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// CreateBudgetRule creates a new Agent Network budget rule +func (a *AgentNetworkAPI) CreateBudgetRule(ctx context.Context, request api.PostApiAgentNetworkBudgetRulesJSONRequestBody) (*api.AgentNetworkBudgetRule, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/budget-rules", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// UpdateBudgetRule updates an Agent Network budget rule +func (a *AgentNetworkAPI) UpdateBudgetRule(ctx context.Context, ruleID string, request api.PutApiAgentNetworkBudgetRulesRuleIdJSONRequestBody) (*api.AgentNetworkBudgetRule, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/budget-rules/"+ruleID, bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkBudgetRule](resp) + return &ret, err +} + +// DeleteBudgetRule deletes an Agent Network budget rule +func (a *AgentNetworkAPI) DeleteBudgetRule(ctx context.Context, ruleID string) error { + resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/budget-rules/"+ruleID, nil, nil) + if err != nil { + return err + } + if resp.Body != nil { + defer resp.Body.Close() + } + + return nil +} + +// GetSettings gets the account's Agent Network gateway settings (cluster, +// subdomain, endpoint, collection toggles). Returns an APIError with +// StatusCode 404 (matchable via IsNotFound) when the account has not been +// bootstrapped yet — bootstrap via UpdateSettings with a cluster, or by +// creating the first provider with bootstrap_cluster set. Management servers +// prior to the 404 contract answered 200 with a JSON null body in that case; +// that legacy shape is translated to the same 404 APIError here so callers +// only ever branch on IsNotFound. +func (a *AgentNetworkAPI) GetSettings(ctx context.Context) (*api.AgentNetworkSettings, error) { + resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/settings", nil, nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, err + } + if trimmed := bytes.TrimSpace(body); len(trimmed) == 0 || bytes.Equal(trimmed, []byte("null")) { + return nil, &APIError{StatusCode: http.StatusNotFound, Message: "agent network settings not found"} + } + var ret api.AgentNetworkSettings + if err := json.Unmarshal(body, &ret); err != nil { + return nil, err + } + return &ret, nil +} + +// UpdateSettings updates the account's Agent Network settings; the request +// replaces every mutable field (collection toggles and retention). Setting +// request.Cluster bootstraps the settings row when the account does not have +// one yet; on a bootstrapped account it must match the assigned cluster (or +// be nil) and any other value is rejected — the cluster is immutable. +func (a *AgentNetworkAPI) UpdateSettings(ctx context.Context, request api.PutApiAgentNetworkSettingsJSONRequestBody) (*api.AgentNetworkSettings, error) { + requestBytes, err := json.Marshal(request) + if err != nil { + return nil, err + } + resp, err := a.c.NewRequest(ctx, "PUT", "/api/agent-network/settings", bytes.NewReader(requestBytes), nil) + if err != nil { + return nil, err + } + if resp.Body != nil { + defer resp.Body.Close() + } + ret, err := parseResponse[api.AgentNetworkSettings](resp) + return &ret, err +} diff --git a/shared/management/client/rest/agentnetwork_test.go b/shared/management/client/rest/agentnetwork_test.go new file mode 100644 index 000000000..a5ad001e2 --- /dev/null +++ b/shared/management/client/rest/agentnetwork_test.go @@ -0,0 +1,476 @@ +//go:build integration + +package rest_test + +import ( + "context" + "encoding/json" + "io" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/client/rest" + "github.com/netbirdio/netbird/shared/management/http/api" + "github.com/netbirdio/netbird/shared/management/http/util" +) + +var ( + testAgentNetworkProvider = api.AgentNetworkProvider{ + Id: "ainp_test", + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + Models: []api.AgentNetworkProviderModel{}, + Enabled: true, + } + + testAgentNetworkPolicy = api.AgentNetworkPolicy{ + Id: "ainpol_test", + Name: "Engineering → OpenAI", + Enabled: true, + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + } + + testAgentNetworkGuardrail = api.AgentNetworkGuardrail{ + Id: "aingr_test", + Name: "No secrets", + } + + testAgentNetworkBudgetRule = api.AgentNetworkBudgetRule{ + Id: "ainbud_test", + Name: "Org monthly ceiling", + Enabled: true, + } + + testAgentNetworkSettings = api.AgentNetworkSettings{ + Cluster: "eu.proxy.netbird.io", + Subdomain: "violet", + Endpoint: "violet.eu.proxy.netbird.io", + EnableLogCollection: true, + AccessLogRetentionDays: ptr(30), + } +) + +func TestAgentNetwork_ListCatalogProviders_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/catalog/providers", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkCatalogProvider{{Id: "openai_api", Name: "OpenAI"}}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListCatalogProviders(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, "openai_api", ret[0].Id) + }) +} + +func TestAgentNetwork_ListProviders_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkProvider{testAgentNetworkProvider}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListProviders(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkProvider, ret[0]) + }) +} + +func TestAgentNetwork_GetProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "GET", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetProvider(context.Background(), "ainp_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_GetProvider_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "not found", Code: 404}) + w.WriteHeader(404) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.GetProvider(context.Background(), "ainp_test") + require.Error(t, err) + assert.True(t, rest.IsNotFound(err), "a 404 must be matchable via IsNotFound") + }) +} + +func TestAgentNetwork_CreateProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + var req api.PostApiAgentNetworkProvidersJSONRequestBody + require.NoError(t, json.Unmarshal(reqBytes, &req)) + assert.Equal(t, "OpenAI", req.Name) + require.NotNil(t, req.BootstrapCluster) + assert.Equal(t, "eu.proxy.netbird.io", *req.BootstrapCluster) + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateProvider(context.Background(), api.PostApiAgentNetworkProvidersJSONRequestBody{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + ApiKey: ptr("sk-test"), + BootstrapCluster: ptr("eu.proxy.netbird.io"), + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_UpdateProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + // Omitted optional fields must be absent from the wire (not + // zero-valued) so the server-side merge preserves them. + assert.NotContains(t, string(reqBytes), "api_key") + assert.NotContains(t, string(reqBytes), "models") + retBytes, _ := json.Marshal(testAgentNetworkProvider) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateProvider(context.Background(), "ainp_test", api.PutApiAgentNetworkProvidersProviderIdJSONRequestBody{ + ProviderId: "openai_api", + Name: "OpenAI", + UpstreamUrl: "https://api.openai.com", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkProvider, *ret) + }) +} + +func TestAgentNetwork_DeleteProvider_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/providers/ainp_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteProvider(context.Background(), "ainp_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListPolicies_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkPolicy{testAgentNetworkPolicy}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListPolicies(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkPolicy, ret[0]) + }) +} + +func TestAgentNetwork_GetPolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetPolicy(context.Background(), "ainpol_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_CreatePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreatePolicy(context.Background(), api.PostApiAgentNetworkPoliciesJSONRequestBody{ + Name: "Engineering → OpenAI", + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_UpdatePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkPolicy) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdatePolicy(context.Background(), "ainpol_test", api.PutApiAgentNetworkPoliciesPolicyIdJSONRequestBody{ + Name: "Engineering → OpenAI", + SourceGroups: []string{"grp-eng"}, + DestinationProviderIds: []string{"ainp_test"}, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkPolicy, *ret) + }) +} + +func TestAgentNetwork_DeletePolicy_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/policies/ainpol_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeletePolicy(context.Background(), "ainpol_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListGuardrails_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkGuardrail{testAgentNetworkGuardrail}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListGuardrails(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkGuardrail, ret[0]) + }) +} + +func TestAgentNetwork_GetGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetGuardrail(context.Background(), "aingr_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_CreateGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateGuardrail(context.Background(), api.PostApiAgentNetworkGuardrailsJSONRequestBody{ + Name: "No secrets", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_UpdateGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkGuardrail) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateGuardrail(context.Background(), "aingr_test", api.PutApiAgentNetworkGuardrailsGuardrailIdJSONRequestBody{ + Name: "No secrets", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkGuardrail, *ret) + }) +} + +func TestAgentNetwork_DeleteGuardrail_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/guardrails/aingr_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteGuardrail(context.Background(), "aingr_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_ListBudgetRules_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal([]api.AgentNetworkBudgetRule{testAgentNetworkBudgetRule}) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.ListBudgetRules(context.Background()) + require.NoError(t, err) + assert.Len(t, ret, 1) + assert.Equal(t, testAgentNetworkBudgetRule, ret[0]) + }) +} + +func TestAgentNetwork_GetBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetBudgetRule(context.Background(), "ainbud_test") + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_CreateBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "POST", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.CreateBudgetRule(context.Background(), api.PostApiAgentNetworkBudgetRulesJSONRequestBody{ + Name: "Org monthly ceiling", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_UpdateBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + retBytes, _ := json.Marshal(testAgentNetworkBudgetRule) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateBudgetRule(context.Background(), "ainbud_test", api.PutApiAgentNetworkBudgetRulesRuleIdJSONRequestBody{ + Name: "Org monthly ceiling", + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkBudgetRule, *ret) + }) +} + +func TestAgentNetwork_DeleteBudgetRule_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/budget-rules/ainbud_test", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "DELETE", r.Method) + _, err := w.Write([]byte("{}")) + require.NoError(t, err) + }) + err := c.AgentNetwork.DeleteBudgetRule(context.Background(), "ainbud_test") + require.NoError(t, err) + }) +} + +func TestAgentNetwork_GetSettings_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(testAgentNetworkSettings) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkSettings, *ret) + }) +} + +func TestAgentNetwork_GetSettings_404(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "agent network settings not found", Code: 404}) + w.WriteHeader(404) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.GetSettings(context.Background()) + require.Error(t, err) + assert.True(t, rest.IsNotFound(err), "unbootstrapped settings must be matchable via IsNotFound") + }) +} + +// TestAgentNetwork_GetSettings_LegacyNullBody pins the compatibility shim for +// management servers that answered 200 with a JSON null body before the 404 +// contract: the client must translate that shape into the same IsNotFound +// error instead of returning a bogus zero-valued settings object. +func TestAgentNetwork_GetSettings_LegacyNullBody(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + _, err := w.Write([]byte("null")) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.GetSettings(context.Background()) + require.Error(t, err) + assert.Nil(t, ret) + assert.True(t, rest.IsNotFound(err), "the legacy 200+null shape must surface as IsNotFound") + }) +} + +func TestAgentNetwork_UpdateSettings_200(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "PUT", r.Method) + reqBytes, err := io.ReadAll(r.Body) + require.NoError(t, err) + var req api.PutApiAgentNetworkSettingsJSONRequestBody + require.NoError(t, json.Unmarshal(reqBytes, &req)) + require.NotNil(t, req.Cluster, "bootstrap cluster must be on the wire") + assert.Equal(t, "eu.proxy.netbird.io", *req.Cluster) + assert.True(t, req.EnableLogCollection) + retBytes, _ := json.Marshal(testAgentNetworkSettings) + _, err = w.Write(retBytes) + require.NoError(t, err) + }) + ret, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{ + Cluster: ptr("eu.proxy.netbird.io"), + EnableLogCollection: true, + }) + require.NoError(t, err) + assert.Equal(t, testAgentNetworkSettings, *ret) + }) +} + +func TestAgentNetwork_UpdateSettings_Err(t *testing.T) { + withMockClient(func(c *rest.Client, mux *http.ServeMux) { + mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) { + retBytes, _ := json.Marshal(util.ErrorResponse{Message: "cluster is immutable once assigned (current: eu.proxy.netbird.io)", Code: 422}) + w.WriteHeader(422) + _, err := w.Write(retBytes) + require.NoError(t, err) + }) + _, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{ + Cluster: ptr("us.proxy.netbird.io"), + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "immutable") + }) +} diff --git a/shared/management/client/rest/client.go b/shared/management/client/rest/client.go index 43312b9e6..6154a6637 100644 --- a/shared/management/client/rest/client.go +++ b/shared/management/client/rest/client.go @@ -147,6 +147,10 @@ type Client struct { // ReverseProxyTokens account-scoped proxy access tokens used to register // self-hosted (bring-your-own-proxy) `netbird proxy` instances. ReverseProxyTokens *ReverseProxyTokensAPI + + // AgentNetwork NetBird Agent Network (AI/LLM gateway) APIs: catalog, + // providers, policies, guardrails, budget rules and account settings. + AgentNetwork *AgentNetworkAPI } // New initialize new Client instance using PAT token @@ -209,6 +213,7 @@ func (c *Client) initialize() { c.ReverseProxyClusters = &ReverseProxyClustersAPI{c} c.ReverseProxyDomains = &ReverseProxyDomainsAPI{c} c.ReverseProxyTokens = &ReverseProxyTokensAPI{c} + c.AgentNetwork = &AgentNetworkAPI{c} } // NewRequest creates and executes new management API request