mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-13 10:19:07 +02:00
The bootstrap check refused pinning onto a cluster another account runs, but nothing held that decision afterwards: a proxy registration consulted the proxies table alone, so the same host could be claimed by another account a moment later, or a week later, and the pin it stranded was immutable. Validating only at bootstrap meant the check was true when it ran and not after. Make the claim symmetric. IsClusterAddressAvailable now treats a gateway pin as what it is — a claim on the host, served by whichever proxy declares that address — and refuses a proxy from a different account, since an account-scoped proxy never receives another account's mappings and so cannot serve the pin it would displace. Both claims are checked in one place so a caller cannot consult one and forget the other. An account claiming the address its own gateway is pinned to is the documented order, not a conflict: pin first, deploy the proxy after. Shared (NetBird-operated) proxies register without an account and are unaffected; an account-scoped proxy reaching for a shared cluster address that accounts are pinned to was already refused by the proxy-row conflict and now stays refused even if those rows are momentarily gone. Whichever of the two lands first now wins and the second is refused, which is what the bootstrap check could not do on its own. A genuinely concurrent pair can still pass both checks — closing that needs an invariant spanning the proxies and settings tables, not a wider read.
750 lines
29 KiB
Go
750 lines
29 KiB
Go
package store
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"math"
|
|
"time"
|
|
|
|
log "github.com/sirupsen/logrus"
|
|
"gorm.io/gorm"
|
|
"gorm.io/gorm/clause"
|
|
|
|
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
|
"github.com/netbirdio/netbird/shared/management/status"
|
|
)
|
|
|
|
// GetAllAgentNetworkProviders returns Agent Network providers across
|
|
// every account. Used by the synthesizer to build the global service map.
|
|
func (s *SqlStore) GetAllAgentNetworkProviders(ctx context.Context, lockStrength LockingStrength) ([]*agentNetworkTypes.Provider, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var providers []*agentNetworkTypes.Provider
|
|
if result := tx.Find(&providers); result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to get all agent network providers from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get all agent network providers from store")
|
|
}
|
|
|
|
for _, provider := range providers {
|
|
if err := provider.DecryptSensitiveData(s.fieldEncrypt); err != nil {
|
|
log.WithContext(ctx).Errorf("failed to decrypt agent network provider %s: %v", provider.ID, err)
|
|
return nil, status.Errorf(status.Internal, "failed to decrypt agent network provider")
|
|
}
|
|
}
|
|
|
|
return providers, nil
|
|
}
|
|
|
|
// GetAgentNetworkMetrics returns aggregated agent-network adoption + usage
|
|
// counts for the self-hosted metrics worker. Each value is a single cheap
|
|
// aggregate; token/cost are summed over the always-collected per-request usage
|
|
// ledger (independent of the log-collection toggle) so they reflect real usage.
|
|
func (s *SqlStore) GetAgentNetworkMetrics(ctx context.Context) (AgentNetworkMetrics, error) {
|
|
var m AgentNetworkMetrics
|
|
db := s.db.WithContext(ctx)
|
|
|
|
// Providers + distinct adopting accounts in one round-trip.
|
|
provRow := db.Model(&agentNetworkTypes.Provider{}).
|
|
Select("COUNT(*) AS providers, COUNT(DISTINCT account_id) AS accounts").Row()
|
|
if err := provRow.Scan(&m.Providers, &m.Accounts); err != nil {
|
|
return AgentNetworkMetrics{}, fmt.Errorf("scan agent network provider metrics: %w", err)
|
|
}
|
|
|
|
if err := db.Model(&agentNetworkTypes.Policy{}).Count(&m.Policies).Error; err != nil {
|
|
return AgentNetworkMetrics{}, fmt.Errorf("count agent network policies: %w", err)
|
|
}
|
|
|
|
if err := db.Model(&agentNetworkTypes.AccountBudgetRule{}).Count(&m.BudgetRules).Error; err != nil {
|
|
return AgentNetworkMetrics{}, fmt.Errorf("count agent network budget rules: %w", err)
|
|
}
|
|
|
|
if err := db.Model(&agentNetworkTypes.Settings{}).
|
|
Where("enable_log_collection = ?", true).Count(&m.LogCollectionEnabled).Error; err != nil {
|
|
return AgentNetworkMetrics{}, fmt.Errorf("count agent network log-collection accounts: %w", err)
|
|
}
|
|
|
|
// COALESCE so an empty ledger scans as 0 instead of NULL.
|
|
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()
|
|
if err := usageRow.Scan(&m.InputTokens, &m.OutputTokens, &m.CostUSD); err != nil {
|
|
return AgentNetworkMetrics{}, fmt.Errorf("scan agent network usage metrics: %w", err)
|
|
}
|
|
|
|
return m, nil
|
|
}
|
|
|
|
func (s *SqlStore) GetAccountAgentNetworkProviders(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*agentNetworkTypes.Provider, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var providers []*agentNetworkTypes.Provider
|
|
result := tx.Find(&providers, accountIDCondition, accountID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to get agent network providers from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network providers from store")
|
|
}
|
|
|
|
for _, provider := range providers {
|
|
if err := provider.DecryptSensitiveData(s.fieldEncrypt); err != nil {
|
|
log.WithContext(ctx).Errorf("failed to decrypt agent network provider %s: %v", provider.ID, err)
|
|
return nil, status.Errorf(status.Internal, "failed to decrypt agent network provider")
|
|
}
|
|
}
|
|
|
|
return providers, nil
|
|
}
|
|
|
|
func (s *SqlStore) GetAgentNetworkProviderByID(ctx context.Context, lockStrength LockingStrength, accountID, providerID string) (*agentNetworkTypes.Provider, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var provider *agentNetworkTypes.Provider
|
|
result := tx.Take(&provider, accountAndIDQueryCondition, accountID, providerID)
|
|
if result.Error != nil {
|
|
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
|
return nil, status.NewAgentNetworkProviderNotFoundError(providerID)
|
|
}
|
|
|
|
log.WithContext(ctx).Errorf("failed to get agent network provider from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network provider from store")
|
|
}
|
|
|
|
if err := provider.DecryptSensitiveData(s.fieldEncrypt); err != nil {
|
|
log.WithContext(ctx).Errorf("failed to decrypt agent network provider %s: %v", provider.ID, err)
|
|
return nil, status.Errorf(status.Internal, "failed to decrypt agent network provider")
|
|
}
|
|
|
|
return provider, nil
|
|
}
|
|
|
|
func (s *SqlStore) SaveAgentNetworkProvider(ctx context.Context, provider *agentNetworkTypes.Provider) error {
|
|
providerCopy := provider.Copy()
|
|
if err := providerCopy.EncryptSensitiveData(s.fieldEncrypt); err != nil {
|
|
log.WithContext(ctx).Errorf("failed to encrypt agent network provider %s: %v", provider.ID, err)
|
|
return status.Errorf(status.Internal, "failed to encrypt agent network provider")
|
|
}
|
|
|
|
result := s.db.Save(providerCopy)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to save agent network provider to store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to save agent network provider to store")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *SqlStore) DeleteAgentNetworkProvider(ctx context.Context, accountID, providerID string) error {
|
|
result := s.db.Delete(&agentNetworkTypes.Provider{}, accountAndIDQueryCondition, accountID, providerID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to delete agent network provider from store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to delete agent network provider from store")
|
|
}
|
|
|
|
if result.RowsAffected == 0 {
|
|
return status.NewAgentNetworkProviderNotFoundError(providerID)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *SqlStore) GetAccountAgentNetworkPolicies(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*agentNetworkTypes.Policy, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var policies []*agentNetworkTypes.Policy
|
|
result := tx.Find(&policies, accountIDCondition, accountID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to get agent network policies from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network policies from store")
|
|
}
|
|
|
|
return policies, nil
|
|
}
|
|
|
|
func (s *SqlStore) GetAgentNetworkPolicyByID(ctx context.Context, lockStrength LockingStrength, accountID, policyID string) (*agentNetworkTypes.Policy, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var policy *agentNetworkTypes.Policy
|
|
result := tx.Take(&policy, accountAndIDQueryCondition, accountID, policyID)
|
|
if result.Error != nil {
|
|
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
|
return nil, status.NewAgentNetworkPolicyNotFoundError(policyID)
|
|
}
|
|
|
|
log.WithContext(ctx).Errorf("failed to get agent network policy from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network policy from store")
|
|
}
|
|
|
|
return policy, nil
|
|
}
|
|
|
|
func (s *SqlStore) SaveAgentNetworkPolicy(ctx context.Context, policy *agentNetworkTypes.Policy) error {
|
|
result := s.db.Save(policy)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to save agent network policy to store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to save agent network policy to store")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *SqlStore) DeleteAgentNetworkPolicy(ctx context.Context, accountID, policyID string) error {
|
|
result := s.db.Delete(&agentNetworkTypes.Policy{}, accountAndIDQueryCondition, accountID, policyID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to delete agent network policy from store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to delete agent network policy from store")
|
|
}
|
|
|
|
if result.RowsAffected == 0 {
|
|
return status.NewAgentNetworkPolicyNotFoundError(policyID)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *SqlStore) GetAccountAgentNetworkGuardrails(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*agentNetworkTypes.Guardrail, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var guardrails []*agentNetworkTypes.Guardrail
|
|
result := tx.Find(&guardrails, accountIDCondition, accountID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to get agent network guardrails from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network guardrails from store")
|
|
}
|
|
|
|
return guardrails, nil
|
|
}
|
|
|
|
func (s *SqlStore) GetAgentNetworkGuardrailByID(ctx context.Context, lockStrength LockingStrength, accountID, guardrailID string) (*agentNetworkTypes.Guardrail, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var guardrail *agentNetworkTypes.Guardrail
|
|
result := tx.Take(&guardrail, accountAndIDQueryCondition, accountID, guardrailID)
|
|
if result.Error != nil {
|
|
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
|
return nil, status.NewAgentNetworkGuardrailNotFoundError(guardrailID)
|
|
}
|
|
|
|
log.WithContext(ctx).Errorf("failed to get agent network guardrail from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network guardrail from store")
|
|
}
|
|
|
|
return guardrail, nil
|
|
}
|
|
|
|
func (s *SqlStore) SaveAgentNetworkGuardrail(ctx context.Context, guardrail *agentNetworkTypes.Guardrail) error {
|
|
result := s.db.Save(guardrail)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to save agent network guardrail to store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to save agent network guardrail to store")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func (s *SqlStore) DeleteAgentNetworkGuardrail(ctx context.Context, accountID, guardrailID string) error {
|
|
result := s.db.Delete(&agentNetworkTypes.Guardrail{}, accountAndIDQueryCondition, accountID, guardrailID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to delete agent network guardrail from store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to delete agent network guardrail from store")
|
|
}
|
|
|
|
if result.RowsAffected == 0 {
|
|
return status.NewAgentNetworkGuardrailNotFoundError(guardrailID)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// GetAgentNetworkSettings returns the per-account Agent Network
|
|
// settings row. Returns status.NotFound when no row exists.
|
|
func (s *SqlStore) GetAgentNetworkSettings(ctx context.Context, lockStrength LockingStrength, accountID string) (*agentNetworkTypes.Settings, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var settings agentNetworkTypes.Settings
|
|
result := tx.Take(&settings, "account_id = ?", accountID)
|
|
if result.Error != nil {
|
|
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
|
return nil, status.Errorf(status.NotFound, "agent network settings for account %s not found", accountID)
|
|
}
|
|
|
|
log.WithContext(ctx).Errorf("failed to get agent network settings from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network settings from store")
|
|
}
|
|
|
|
return &settings, nil
|
|
}
|
|
|
|
// GetAllAgentNetworkSettings returns every account's settings row. Used by the
|
|
// access-log retention sweep to learn each account's retention window.
|
|
func (s *SqlStore) GetAllAgentNetworkSettings(ctx context.Context, lockStrength LockingStrength) ([]*agentNetworkTypes.Settings, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var settings []*agentNetworkTypes.Settings
|
|
if err := tx.Find(&settings).Error; err != nil {
|
|
log.WithContext(ctx).Errorf("failed to list agent network settings: %v", err)
|
|
return nil, status.Errorf(status.Internal, "failed to list agent network settings")
|
|
}
|
|
return settings, nil
|
|
}
|
|
|
|
// HasGatewayPinnedByOtherAccount reports whether an account other than the
|
|
// given one has its agent network gateway pinned to this host.
|
|
//
|
|
// A pin is a claim on the host, the same way a proxy row is: the pinned
|
|
// endpoint is served by whichever proxy declares that address, and an
|
|
// account-scoped proxy only ever receives its own account's mappings. A proxy
|
|
// from a different account taking the address therefore cannot serve the pin
|
|
// and silently strands it. The pin is immutable, so the account that holds it
|
|
// cannot move out of the way — the later claimant is the one to refuse.
|
|
//
|
|
// Both sides are canonical (settings normalize on write, proxy addresses
|
|
// canonicalize at connect), so the match is exact and uses the proxy_address
|
|
// index.
|
|
func (s *SqlStore) HasGatewayPinnedByOtherAccount(ctx context.Context, host, accountID string) (bool, error) {
|
|
var count int64
|
|
result := s.db.
|
|
Model(&agentNetworkTypes.Settings{}).
|
|
Where("proxy_address = ? AND account_id != ?", host, accountID).
|
|
Count(&count)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to check agent network gateway pins by proxy address: %v", result.Error)
|
|
return false, status.Errorf(status.Internal, "check agent network gateway pins")
|
|
}
|
|
return count > 0, nil
|
|
}
|
|
|
|
// GetAgentNetworkSettingsByProxyAddress returns every Settings row whose
|
|
// gateway is served by the proxy declaring the given cluster address. Used by
|
|
// cluster-scoped synthesis to find the accounts a shared proxy serves.
|
|
func (s *SqlStore) GetAgentNetworkSettingsByProxyAddress(ctx context.Context, lockStrength LockingStrength, proxyAddress string) ([]*agentNetworkTypes.Settings, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var settings []*agentNetworkTypes.Settings
|
|
result := tx.Find(&settings, "proxy_address = ?", proxyAddress)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to get agent network settings by proxy address from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network settings by proxy address from store")
|
|
}
|
|
|
|
return settings, nil
|
|
}
|
|
|
|
// GetAgentNetworkSettingsByDomain resolves the single Settings row holding the
|
|
// given endpoint hostname — a point query on the domain unique index. Returns
|
|
// status.NotFound when no account owns the domain.
|
|
func (s *SqlStore) GetAgentNetworkSettingsByDomain(ctx context.Context, lockStrength LockingStrength, domain string) (*agentNetworkTypes.Settings, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var settings agentNetworkTypes.Settings
|
|
result := tx.Take(&settings, "domain = ?", domain)
|
|
if result.Error != nil {
|
|
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
|
return nil, status.Errorf(status.NotFound, "agent network settings for domain %s not found", domain)
|
|
}
|
|
|
|
log.WithContext(ctx).Errorf("failed to get agent network settings by domain from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network settings by domain from store")
|
|
}
|
|
|
|
return &settings, nil
|
|
}
|
|
|
|
// CreateAgentNetworkSettings inserts a new settings row.
|
|
//
|
|
// Unlike SaveAgentNetworkSettings (an upsert) this is a plain INSERT, and it
|
|
// returns the driver error unwrapped. Both properties are required by the
|
|
// bootstrap allocator: an upsert would overwrite whichever row it collided
|
|
// with, and the allocator classifies the rejection by matching the driver's
|
|
// message — a unique violation on the account primary key means a concurrent
|
|
// bootstrap for the same account won, and one on the domain index means the
|
|
// hostname is taken.
|
|
func (s *SqlStore) CreateAgentNetworkSettings(ctx context.Context, settings *agentNetworkTypes.Settings) error {
|
|
if err := s.db.Create(settings).Error; err != nil {
|
|
log.WithContext(ctx).Debugf("failed to create agent network settings: %v", err)
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// SaveAgentNetworkSettings upserts the per-account Agent Network
|
|
// settings row.
|
|
func (s *SqlStore) SaveAgentNetworkSettings(ctx context.Context, settings *agentNetworkTypes.Settings) error {
|
|
result := s.db.Save(settings)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to save agent network settings to store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to save agent network settings to store")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// DeleteAgentNetworkSettings removes the per-account Agent Network settings
|
|
// row, releasing the account's endpoint. Returns status.NotFound when no row
|
|
// exists. The guards on the delete (no providers, no proxy actively serving
|
|
// the endpoint) live in the manager, which runs this inside a transaction
|
|
// after re-checking them under a row lock.
|
|
func (s *SqlStore) DeleteAgentNetworkSettings(ctx context.Context, accountID string) error {
|
|
result := s.db.Delete(&agentNetworkTypes.Settings{}, "account_id = ?", accountID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to delete agent network settings from store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to delete agent network settings from store")
|
|
}
|
|
|
|
if result.RowsAffected == 0 {
|
|
return status.Errorf(status.NotFound, "agent network settings for account %s not found", accountID)
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// IncrementAgentNetworkConsumption atomically upserts the consumption
|
|
// row keyed on (account, dim_kind, dim_id, window_seconds, window_start)
|
|
// and adds the supplied deltas. Concurrent calls from multiple proxy
|
|
// nodes converge — the database performs the increment server-side via
|
|
// ON CONFLICT DO UPDATE so no read-modify-write race exists.
|
|
func (s *SqlStore) IncrementAgentNetworkConsumption(
|
|
ctx context.Context,
|
|
accountID string,
|
|
kind agentNetworkTypes.ConsumptionDimension,
|
|
dimID string,
|
|
windowSeconds int64,
|
|
windowStart time.Time,
|
|
tokensIn, tokensOut int64,
|
|
costUSD float64,
|
|
) error {
|
|
if accountID == "" || dimID == "" || windowSeconds <= 0 {
|
|
return status.Errorf(status.InvalidArgument, "account_id, dim_id and window_seconds must be set")
|
|
}
|
|
// Deltas are added server-side via ON CONFLICT; a negative or non-finite
|
|
// value would silently decrement / poison the persisted totals.
|
|
if tokensIn < 0 || tokensOut < 0 || costUSD < 0 || math.IsNaN(costUSD) || math.IsInf(costUSD, 0) {
|
|
return status.Errorf(status.InvalidArgument, "consumption deltas must be non-negative and finite")
|
|
}
|
|
row := agentNetworkTypes.Consumption{
|
|
AccountID: accountID,
|
|
DimensionKind: kind,
|
|
DimensionID: dimID,
|
|
WindowSeconds: windowSeconds,
|
|
WindowStartUTC: windowStart.UTC(),
|
|
TokensInput: tokensIn,
|
|
TokensOutput: tokensOut,
|
|
CostUSD: costUSD,
|
|
UpdatedAt: time.Now().UTC(),
|
|
}
|
|
const tbl = "agent_network_consumption"
|
|
err := s.db.Clauses(clause.OnConflict{
|
|
Columns: []clause.Column{
|
|
{Name: "account_id"},
|
|
{Name: "dim_kind"},
|
|
{Name: "dim_id"},
|
|
{Name: "window_seconds"},
|
|
{Name: "window_start_utc"},
|
|
},
|
|
DoUpdates: clause.Assignments(map[string]any{
|
|
"tokens_input": gorm.Expr(tbl+".tokens_input + ?", tokensIn),
|
|
"tokens_output": gorm.Expr(tbl+".tokens_output + ?", tokensOut),
|
|
"cost_usd": gorm.Expr(tbl+".cost_usd + ?", costUSD),
|
|
"updated_at": time.Now().UTC(),
|
|
}),
|
|
}).Create(&row).Error
|
|
if err != nil {
|
|
log.WithContext(ctx).Errorf("failed to increment agent network consumption: %v", err)
|
|
return status.Errorf(status.Internal, "failed to increment agent network consumption")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// GetAgentNetworkConsumption returns the consumption row for the exact
|
|
// window key. Returns a zero-valued row (not found mapped to zero) so
|
|
// callers can use the result as the headroom basis without nil checks.
|
|
func (s *SqlStore) GetAgentNetworkConsumption(
|
|
ctx context.Context,
|
|
lockStrength LockingStrength,
|
|
accountID string,
|
|
kind agentNetworkTypes.ConsumptionDimension,
|
|
dimID string,
|
|
windowSeconds int64,
|
|
windowStart time.Time,
|
|
) (*agentNetworkTypes.Consumption, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
var row agentNetworkTypes.Consumption
|
|
result := tx.Take(&row,
|
|
"account_id = ? AND dim_kind = ? AND dim_id = ? AND window_seconds = ? AND window_start_utc = ?",
|
|
accountID, kind, dimID, windowSeconds, windowStart.UTC())
|
|
if result.Error != nil {
|
|
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
|
return &agentNetworkTypes.Consumption{
|
|
AccountID: accountID,
|
|
DimensionKind: kind,
|
|
DimensionID: dimID,
|
|
WindowSeconds: windowSeconds,
|
|
WindowStartUTC: windowStart.UTC(),
|
|
}, nil
|
|
}
|
|
log.WithContext(ctx).Errorf("failed to get agent network consumption: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network consumption")
|
|
}
|
|
return &row, nil
|
|
}
|
|
|
|
// GetAgentNetworkConsumptionBatch reads many consumption counters for one
|
|
// account in a single query, returning a map keyed by the exact
|
|
// ConsumptionKey. Missing counters are simply absent from the map (callers
|
|
// treat absence as a zero counter). Replaces the per-cap point reads the
|
|
// policy selector previously issued one at a time.
|
|
func (s *SqlStore) GetAgentNetworkConsumptionBatch(
|
|
ctx context.Context,
|
|
lockStrength LockingStrength,
|
|
accountID string,
|
|
keys []agentNetworkTypes.ConsumptionKey,
|
|
) (map[agentNetworkTypes.ConsumptionKey]*agentNetworkTypes.Consumption, error) {
|
|
out := make(map[agentNetworkTypes.ConsumptionKey]*agentNetworkTypes.Consumption, len(keys))
|
|
if len(keys) == 0 {
|
|
return out, nil
|
|
}
|
|
|
|
// Collect the distinct dim ids, windows and window starts so a single
|
|
// query scopes to exactly the current windows in play, then filter the
|
|
// returned rows down to the exact requested keys.
|
|
wanted := make(map[agentNetworkTypes.ConsumptionKey]struct{}, len(keys))
|
|
dimSet := make(map[string]struct{})
|
|
winSet := make(map[int64]struct{})
|
|
startSet := make(map[time.Time]struct{})
|
|
for _, k := range keys {
|
|
k.WindowStartUTC = k.WindowStartUTC.UTC()
|
|
wanted[k] = struct{}{}
|
|
dimSet[k.DimID] = struct{}{}
|
|
winSet[k.WindowSeconds] = struct{}{}
|
|
startSet[k.WindowStartUTC] = struct{}{}
|
|
}
|
|
dimIDs := make([]string, 0, len(dimSet))
|
|
for d := range dimSet {
|
|
dimIDs = append(dimIDs, d)
|
|
}
|
|
windows := make([]int64, 0, len(winSet))
|
|
for w := range winSet {
|
|
windows = append(windows, w)
|
|
}
|
|
starts := make([]time.Time, 0, len(startSet))
|
|
for t := range startSet {
|
|
starts = append(starts, t)
|
|
}
|
|
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
var rows []*agentNetworkTypes.Consumption
|
|
result := tx.Find(&rows,
|
|
"account_id = ? AND dim_id IN ? AND window_seconds IN ? AND window_start_utc IN ?",
|
|
accountID, dimIDs, windows, starts)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to batch-get agent network consumption: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network consumption")
|
|
}
|
|
for _, row := range rows {
|
|
k := agentNetworkTypes.ConsumptionKey{
|
|
Kind: row.DimensionKind,
|
|
DimID: row.DimensionID,
|
|
WindowSeconds: row.WindowSeconds,
|
|
WindowStartUTC: row.WindowStartUTC.UTC(),
|
|
}
|
|
if _, ok := wanted[k]; ok {
|
|
out[k] = row
|
|
}
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// IncrementAgentNetworkConsumptionBatch applies the same usage delta to every
|
|
// supplied counter inside a single transaction, so all per-(dimension, window)
|
|
// counters a served request books are written atomically in one round-trip
|
|
// instead of one upsert per counter. Keys are deduplicated by the caller.
|
|
func (s *SqlStore) IncrementAgentNetworkConsumptionBatch(
|
|
ctx context.Context,
|
|
accountID string,
|
|
keys []agentNetworkTypes.ConsumptionKey,
|
|
tokensIn, tokensOut int64,
|
|
costUSD float64,
|
|
) error {
|
|
if accountID == "" {
|
|
return status.Errorf(status.InvalidArgument, "account_id must be set")
|
|
}
|
|
if tokensIn < 0 || tokensOut < 0 || costUSD < 0 || math.IsNaN(costUSD) || math.IsInf(costUSD, 0) {
|
|
return status.Errorf(status.InvalidArgument, "consumption deltas must be non-negative and finite")
|
|
}
|
|
if len(keys) == 0 {
|
|
return nil
|
|
}
|
|
|
|
const tbl = "agent_network_consumption"
|
|
err := s.db.Transaction(func(tx *gorm.DB) error {
|
|
for _, k := range keys {
|
|
if k.DimID == "" || k.WindowSeconds <= 0 {
|
|
return status.Errorf(status.InvalidArgument, "dim_id and window_seconds must be set")
|
|
}
|
|
now := time.Now().UTC()
|
|
row := agentNetworkTypes.Consumption{
|
|
AccountID: accountID,
|
|
DimensionKind: k.Kind,
|
|
DimensionID: k.DimID,
|
|
WindowSeconds: k.WindowSeconds,
|
|
WindowStartUTC: k.WindowStartUTC.UTC(),
|
|
TokensInput: tokensIn,
|
|
TokensOutput: tokensOut,
|
|
CostUSD: costUSD,
|
|
UpdatedAt: now,
|
|
}
|
|
if err := tx.Clauses(clause.OnConflict{
|
|
Columns: []clause.Column{
|
|
{Name: "account_id"},
|
|
{Name: "dim_kind"},
|
|
{Name: "dim_id"},
|
|
{Name: "window_seconds"},
|
|
{Name: "window_start_utc"},
|
|
},
|
|
DoUpdates: clause.Assignments(map[string]any{
|
|
"tokens_input": gorm.Expr(tbl+".tokens_input + ?", tokensIn),
|
|
"tokens_output": gorm.Expr(tbl+".tokens_output + ?", tokensOut),
|
|
"cost_usd": gorm.Expr(tbl+".cost_usd + ?", costUSD),
|
|
"updated_at": now,
|
|
}),
|
|
}).Create(&row).Error; err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
log.WithContext(ctx).Errorf("failed to batch-increment agent network consumption: %v", err)
|
|
return status.Errorf(status.Internal, "failed to increment agent network consumption")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// ListAgentNetworkConsumption returns every consumption row recorded
|
|
// for the account, ordered by window_start descending. Backs the
|
|
// dashboard's basic counter view.
|
|
func (s *SqlStore) ListAgentNetworkConsumption(
|
|
ctx context.Context,
|
|
lockStrength LockingStrength,
|
|
accountID string,
|
|
) ([]*agentNetworkTypes.Consumption, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
var rows []*agentNetworkTypes.Consumption
|
|
result := tx.
|
|
Order("window_start_utc DESC").
|
|
Find(&rows, accountIDCondition, accountID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to list agent network consumption: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to list agent network consumption")
|
|
}
|
|
return rows, nil
|
|
}
|
|
|
|
// GetAccountAgentNetworkBudgetRules returns every account-level budget rule for
|
|
// the account.
|
|
func (s *SqlStore) GetAccountAgentNetworkBudgetRules(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*agentNetworkTypes.AccountBudgetRule, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var rules []*agentNetworkTypes.AccountBudgetRule
|
|
result := tx.Find(&rules, accountIDCondition, accountID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to get agent network budget rules from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network budget rules from store")
|
|
}
|
|
|
|
return rules, nil
|
|
}
|
|
|
|
// GetAgentNetworkBudgetRuleByID returns a single budget rule scoped to the
|
|
// account, or a NotFound error.
|
|
func (s *SqlStore) GetAgentNetworkBudgetRuleByID(ctx context.Context, lockStrength LockingStrength, accountID, ruleID string) (*agentNetworkTypes.AccountBudgetRule, error) {
|
|
tx := s.db
|
|
if lockStrength != LockingStrengthNone {
|
|
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
|
}
|
|
|
|
var rule *agentNetworkTypes.AccountBudgetRule
|
|
result := tx.Take(&rule, accountAndIDQueryCondition, accountID, ruleID)
|
|
if result.Error != nil {
|
|
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
|
return nil, status.NewAgentNetworkBudgetRuleNotFoundError(ruleID)
|
|
}
|
|
|
|
log.WithContext(ctx).Errorf("failed to get agent network budget rule from store: %v", result.Error)
|
|
return nil, status.Errorf(status.Internal, "failed to get agent network budget rule from store")
|
|
}
|
|
|
|
return rule, nil
|
|
}
|
|
|
|
// SaveAgentNetworkBudgetRule upserts a budget rule.
|
|
func (s *SqlStore) SaveAgentNetworkBudgetRule(ctx context.Context, rule *agentNetworkTypes.AccountBudgetRule) error {
|
|
result := s.db.Save(rule)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to save agent network budget rule to store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to save agent network budget rule to store")
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// DeleteAgentNetworkBudgetRule removes a budget rule scoped to the account.
|
|
func (s *SqlStore) DeleteAgentNetworkBudgetRule(ctx context.Context, accountID, ruleID string) error {
|
|
result := s.db.Delete(&agentNetworkTypes.AccountBudgetRule{}, accountAndIDQueryCondition, accountID, ruleID)
|
|
if result.Error != nil {
|
|
log.WithContext(ctx).Errorf("failed to delete agent network budget rule from store: %v", result.Error)
|
|
return status.Errorf(status.Internal, "failed to delete agent network budget rule from store")
|
|
}
|
|
|
|
if result.RowsAffected == 0 {
|
|
return status.NewAgentNetworkBudgetRuleNotFoundError(ruleID)
|
|
}
|
|
|
|
return nil
|
|
}
|