diff --git a/internal/agent/api.go b/internal/agent/api.go index d419d7dac..f3a1b5d99 100644 --- a/internal/agent/api.go +++ b/internal/agent/api.go @@ -1457,6 +1457,33 @@ func resolveOutboundConversationID(explicit, msgType string, referenced *identit return identity.NewConversationID() } +// messageSendReserver is the accept-time reservation surface a limits +// Enforcer may expose (limits.DBEnforcer does). It is asserted rather than +// added to limits.Enforcer so test enforcers that only pre-check stay +// unchanged, and an enforcer without it keeps today's pre-check-only +// behavior. +type messageSendReserver interface { + ReserveMessageSendTx(ctx context.Context, tx pgx.Tx, userID string, units int) error +} + +// limitExceededDetails mirrors httpapi.LimitExceededDetails for a 402 raised +// inside the accept transaction, so both the pre-check and the reservation +// produce the same envelope details. +func limitExceededDetails(le *limits.LimitExceededError) map[string]any { + d := map[string]any{ + "resource": le.Resource, + "limit": le.Limit, + "current": le.Current, + } + if le.Limits.PlanCode != "" { + d["plan_code"] = le.Limits.PlanCode + } + if le.Limits.UpgradeURL != "" { + d["upgrade_url"] = le.Limits.UpgradeURL + } + return d +} + func (a *API) DeliverOutbound(ctx context.Context, user *identity.User, agent *identity.AgentIdentity, req outbound.SendRequest, msgType, replyToEmailMessageID string, referenced *identity.Message, idemCompleteTx AcceptIdemCompleter) (*OutboundResult, *OutboundError) { parentMessageID := "" if msgType == "reply" && referenced != nil { @@ -1592,6 +1619,11 @@ func (a *API) DeliverOutbound(ctx context.Context, user *identity.User, agent *i if scheduledAt != nil { acceptStatus = "scheduled" } + // Units this send costs, counted with the same normalizer the terminal + // meter uses. Only immediate sends reserve here: a scheduled send's flow + // cap is judged against the month it fires in by the worker's fire-time + // gate, not the month it was accepted in. + units := identity.UniqueRecipientCount(comp.To, comp.CC, comp.BCC) var accepted *identity.Message // Crash boundary: // - Before Commit returns, WithTx rolls back the message, River job, and @@ -1602,6 +1634,17 @@ func (a *API) DeliverOutbound(ctx context.Context, user *identity.User, agent *i // 202/message id. Without a key, that ambiguous retry is a new request and // may enqueue a duplicate (the public guarantee is at-least-once). if txErr := a.store.WithTx(ctx, func(tx pgx.Tx) error { + // The accept-time quota reservation runs first, before the insert, so + // the count it reads under the per-user lock cannot miss a send this + // transaction is racing. An immediate send that would cross + // max_messages_month is refused with the message row rolled back. + if scheduledAt == nil { + if reserver, ok := a.enforcer.(messageSendReserver); ok { + if err := reserver.ReserveMessageSendTx(ctx, tx, user.ID, units); err != nil { + return err + } + } + } msg, err := a.store.CreateOutboundMessageThreadedTx(ctx, tx, parentMessageID, agent.ID, comp.To, comp.CC, comp.BCC, req.Subject, msgType, comp.Method, "", req.ConversationID, comp.Raw, "accepted", comp.EnvelopeFrom, comp.SentAs) if err != nil { return err @@ -1632,6 +1675,17 @@ func (a *API) DeliverOutbound(ctx context.Context, user *identity.User, agent *i accepted = msg return nil }); txErr != nil { + if le, ok := limits.IsLimitExceeded(txErr); ok { + // The accept-time reservation refused this send; the message row + // rolled back with it. Carry the same 402 details the pre-check + // path produces so the handler renders one envelope shape. + return nil, &OutboundError{ + Status: http.StatusPaymentRequired, + Code: "limit_exceeded", + Msg: le.Error(), + Details: limitExceededDetails(le), + } + } if errors.Is(txErr, outboundsend.ErrSendingPaused) { // The account is paused for sending abuse: refuse at the door // rather than queue mail that can never leave. Nothing was @@ -1761,6 +1815,14 @@ func (a *API) acceptPlatformSend(ctx context.Context, agent *identity.AgentIdent } var accepted *identity.Message if txErr := a.store.WithTx(ctx, func(tx pgx.Tx) error { + // Same accept-time reservation as DeliverOutbound's immediate branch: + // the platform test send is metered against the owner's flow cap, so + // it must reserve before inserting too. + if reserver, ok := a.enforcer.(messageSendReserver); ok { + if err := reserver.ReserveMessageSendTx(ctx, tx, agent.UserID, identity.UniqueRecipientCount(comp.To, comp.CC, comp.BCC)); err != nil { + return err + } + } msg, err := a.store.CreateOutboundMessageTx(ctx, tx, agent.ID, comp.To, comp.CC, comp.BCC, req.Subject, msgType, comp.Method, "", req.ConversationID, comp.Raw, "accepted", comp.EnvelopeFrom, comp.SentAs) if err != nil { return err @@ -1775,6 +1837,14 @@ func (a *API) acceptPlatformSend(ctx context.Context, agent *identity.AgentIdent accepted = msg return nil }); txErr != nil { + if le, ok := limits.IsLimitExceeded(txErr); ok { + return nil, &OutboundError{ + Status: http.StatusPaymentRequired, + Code: "limit_exceeded", + Msg: le.Error(), + Details: limitExceededDetails(le), + } + } if errors.Is(txErr, outboundsend.ErrSendingPaused) { return nil, &OutboundError{Status: http.StatusForbidden, Code: "sending_paused", Msg: "sending is paused for this account"} } diff --git a/internal/e2e/messages_send_quota_race_e2e_test.go b/internal/e2e/messages_send_quota_race_e2e_test.go new file mode 100644 index 000000000..fdef81c09 --- /dev/null +++ b/internal/e2e/messages_send_quota_race_e2e_test.go @@ -0,0 +1,182 @@ +//go:build integration + +package e2e_test + +import ( + "context" + "fmt" + "io" + "net/http" + "strings" + "sync" + "testing" + + "github.com/tokencanopy/e2a/internal/limits" + "github.com/tokencanopy/e2a/internal/testutil" +) + +// sendMessageURL is the immediate-send endpoint WITHOUT ?wait=sent: these +// tests only care about accept-time admission, and the server is built with +// WithManualJobs so nothing drains the outbound queue. +func sendMessageURL(base, agentEmail string) string { + return base + "/v1/agents/" + agentEmail + "/messages" +} + +// quotaMessageBody is a minimal, valid immediate-send request body. +func quotaMessageBody(i int) string { + return fmt.Sprintf(`{"to":["alice@example.com"],"subject":"quota %d","text":"quota send #%d"}`, i, i) +} + +func postQuotaSend(t *testing.T, url, apiKey string, body string) (int, []byte) { + t.Helper() + req, err := http.NewRequest("POST", url, strings.NewReader(body)) + if err != nil { + t.Fatalf("build request: %v", err) + } + req.Header.Set("Authorization", "Bearer "+apiKey) + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("POST: %v", err) + } + defer resp.Body.Close() + out, _ := io.ReadAll(resp.Body) + return resp.StatusCode, out +} + +// Accepted immediate sends must each consume max_messages_month at accept +// time. The accept-path pre-check reads usage_summaries, which is only written +// once a message terminally sends, so without an accept-time reservation the +// counter stays at 0 for every send in flight and the cap admits an unbounded +// number. This is the sequential boundary: cap 2, three sends, the third must +// be refused. +func TestImmediateSendConsumesMonthlyQuotaAtAccept(t *testing.T) { + pool := testutil.TestDB(t) + ts := testutil.TestServer(t, pool, testutil.WithManualJobs()) + user, key, agent := setupDomainAndAgent(t, ts, "agent@quota.example.com", "quota.example.com", "", "") + if err := limits.NewStore(pool).Upsert(context.Background(), user.ID, limits.Limits{ + PlanCode: "test", MaxAgents: 1000, MaxDomains: 1000, + MaxMessagesMonth: 2, MaxStorageBytes: 1 << 40, + }); err != nil { + t.Fatalf("Upsert limits: %v", err) + } + + url := sendMessageURL(ts.HTTPServer.URL, agent.EmailAddress()) + for i := 1; i <= 2; i++ { + code, body := postQuotaSend(t, url, key.PlaintextKey, quotaMessageBody(i)) + if code != http.StatusAccepted { + t.Fatalf("send %d: status=%d body=%s, want 202 (cap not consumed per send)", i, code, body) + } + } + code, body := postQuotaSend(t, url, key.PlaintextKey, quotaMessageBody(3)) + if code != http.StatusPaymentRequired || !strings.Contains(string(body), `"limit_exceeded"`) { + t.Fatalf("send at the cap: status=%d body=%s, want 402 limit_exceeded", code, body) + } + for _, needle := range []string{`"resource":"messages_month"`, `"limit":2`, `"current":2`} { + if !strings.Contains(string(body), needle) { + t.Errorf("402 body missing %s: %s", needle, body) + } + } +} + +// The per-day cap uses the same accept-time reservation, so it must also count +// a send as soon as it is accepted rather than only once it is metered. +func TestImmediateSendConsumesDailyQuotaAtAccept(t *testing.T) { + pool := testutil.TestDB(t) + ts := testutil.TestServer(t, pool, testutil.WithManualJobs()) + user, key, agent := setupDomainAndAgent(t, ts, "agent@daily.example.com", "daily.example.com", "", "") + one := 1 + if err := limits.NewStore(pool).Upsert(context.Background(), user.ID, limits.Limits{ + PlanCode: "test", MaxAgents: 1000, MaxDomains: 1000, + MaxMessagesMonth: 1000, MaxMessagesDay: &one, MaxStorageBytes: 1 << 40, + }); err != nil { + t.Fatalf("Upsert limits: %v", err) + } + + url := sendMessageURL(ts.HTTPServer.URL, agent.EmailAddress()) + if code, body := postQuotaSend(t, url, key.PlaintextKey, quotaMessageBody(1)); code != http.StatusAccepted { + t.Fatalf("first send: status=%d body=%s, want 202", code, body) + } + code, body := postQuotaSend(t, url, key.PlaintextKey, quotaMessageBody(2)) + if code != http.StatusPaymentRequired || !strings.Contains(string(body), `"resource":"messages_day"`) { + t.Fatalf("second send: status=%d body=%s, want 402 messages_day", code, body) + } +} + +// A burst of concurrent immediate sends must not all pass the same +// pre-increment count. With cap 3 and 10 simultaneous sends exactly 3 may be +// accepted, and exactly 3 accepted rows may exist: the rest must be refused +// with 402 at accept. +func TestConcurrentImmediateSendsRespectMonthlyQuota(t *testing.T) { + pool := testutil.TestDB(t) + ts := testutil.TestServer(t, pool, testutil.WithManualJobs()) + user, key, agent := setupDomainAndAgent(t, ts, "agent@burst.example.com", "burst.example.com", "", "") + if err := limits.NewStore(pool).Upsert(context.Background(), user.ID, limits.Limits{ + PlanCode: "test", MaxAgents: 1000, MaxDomains: 1000, + MaxMessagesMonth: 3, MaxStorageBytes: 1 << 40, + }); err != nil { + t.Fatalf("Upsert limits: %v", err) + } + + const n = 10 + url := sendMessageURL(ts.HTTPServer.URL, agent.EmailAddress()) + codes := make([]int, n) + bodies := make([][]byte, n) + var wg sync.WaitGroup + start := make(chan struct{}) + for i := 0; i < n; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + req, err := http.NewRequest("POST", url, strings.NewReader(quotaMessageBody(i))) + if err != nil { + t.Errorf("build request %d: %v", i, err) + return + } + req.Header.Set("Authorization", "Bearer "+key.PlaintextKey) + req.Header.Set("Content-Type", "application/json") + <-start + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Errorf("POST %d: %v", i, err) + return + } + defer resp.Body.Close() + out, _ := io.ReadAll(resp.Body) + codes[i] = resp.StatusCode + bodies[i] = out + }(i) + } + close(start) + wg.Wait() + + var accepted, rejected int + for i, code := range codes { + switch code { + case http.StatusAccepted: + accepted++ + case http.StatusPaymentRequired: + rejected++ + if !strings.Contains(string(bodies[i]), `"limit_exceeded"`) { + t.Errorf("send %d: 402 body without limit_exceeded: %s", i, bodies[i]) + } + default: + t.Errorf("send %d: unexpected status %d body=%s", i, code, bodies[i]) + } + } + if accepted != 3 || rejected != n-3 { + t.Fatalf("want 3 accepted and %d rejected (max_messages_month=3), got accepted=%d rejected=%d (codes=%v)", + n-3, accepted, rejected, codes) + } + + var acceptedRows int + if err := pool.QueryRow(context.Background(), + `SELECT count(*) FROM messages WHERE agent_id = $1 AND direction = 'outbound' AND delivery_status = 'accepted'`, + agent.ID, + ).Scan(&acceptedRows); err != nil { + t.Fatalf("count accepted messages: %v", err) + } + if acceptedRows != 3 { + t.Fatalf("accepted message rows = %d, want 3 (max_messages_month was oversold)", acceptedRows) + } +} diff --git a/internal/limits/enforcer.go b/internal/limits/enforcer.go index 2ddfedff0..66d4b771c 100644 --- a/internal/limits/enforcer.go +++ b/internal/limits/enforcer.go @@ -5,6 +5,8 @@ import ( "sync" "time" + "github.com/jackc/pgx/v5" + "github.com/tokencanopy/e2a/internal/usage" ) @@ -27,6 +29,20 @@ type limitsReader interface { Get(ctx context.Context, userID string) (Limits, bool, error) } +// counterTxReader is the optional transaction-scoped counter surface used by +// ReserveMessageSendTx. Declared separately from Counter (and asserted, not +// required) so the existing test fakes stay unchanged. +type counterTxReader interface { + MessagesThisMonthTx(ctx context.Context, tx pgx.Tx, userID string) (int, error) + MessagesTodayTx(ctx context.Context, tx pgx.Tx, userID string) (int, error) +} + +// limitsTxReader is the optional transaction-scoped limits surface used by +// ReserveMessageSendTx; same rationale as counterTxReader. +type limitsTxReader interface { + GetTx(ctx context.Context, tx pgx.Tx, userID string) (Limits, bool, error) +} + // DBEnforcer is the production Enforcer: reads account_limits + falls // back to operator Defaults, counts current resources via the usage // store, and caches the resolved Limits in-process for cacheTTL to keep @@ -118,24 +134,28 @@ func (e *DBEnforcer) Get(ctx context.Context, userID string) (Limits, error) { if err != nil { return Limits{}, err } - var resolved Limits - if found { - resolved = row - } else { - resolved = Limits{ - PlanCode: e.defaults.PlanCode, - MaxAgents: e.defaults.MaxAgents, - MaxDomains: e.defaults.MaxDomains, - MaxMessagesMonth: e.defaults.MaxMessagesMonth, - MaxMessagesDay: e.defaults.MaxMessagesDay, - MaxStorageBytes: e.defaults.MaxStorageBytes, - OutboundFooterEnabled: e.defaults.OutboundFooterEnabled, - } - } + resolved := e.resolveLimits(row, found) e.cachePut(userID, resolved, gen) return resolved, nil } +// resolveLimits applies the operator Defaults when the user has no +// account_limits row. +func (e *DBEnforcer) resolveLimits(row Limits, found bool) Limits { + if found { + return row + } + return Limits{ + PlanCode: e.defaults.PlanCode, + MaxAgents: e.defaults.MaxAgents, + MaxDomains: e.defaults.MaxDomains, + MaxMessagesMonth: e.defaults.MaxMessagesMonth, + MaxMessagesDay: e.defaults.MaxMessagesDay, + MaxStorageBytes: e.defaults.MaxStorageBytes, + OutboundFooterEnabled: e.defaults.OutboundFooterEnabled, + } +} + // Invalidate evicts the user's cached Limits and advances the // process-wide invalidation epoch, so any fill already in flight (one // that read account_limits before the caller's write committed) is @@ -290,6 +310,90 @@ func (e *DBEnforcer) CheckMessageSend(ctx context.Context, userID string, units return nil } +// ReserveMessageSendTx is the accept-time, race-proof form of the message-flow +// check: the caller runs it inside the accept transaction, before inserting +// the message, so the counter it reads includes every accepted-but-unmetered +// send already committed. It takes a per-user advisory lock (keyspace 3, +// distinct from claimOrCreateDomain's keyspaces 0/1 and +// CreateAgentWithLimit's keyspace 2) so two simultaneous immediate sends for +// one account cannot both read the same pre-insert total — the same shape as +// the max_agents (#942) and max_domains (#901) fixes. +// +// The month and day caps are read on tx, so the lock never has to wait on a +// second pool connection; the storage stock cap is not re-checked here because +// it is not a flow budget and the accept-time pre-check already enforces it. +func (e *DBEnforcer) ReserveMessageSendTx(ctx context.Context, tx pgx.Tx, userID string, units int) error { + if units < 1 { + units = 1 + } + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1, 3))`, userID); err != nil { + return err + } + lim, err := e.limitsForTx(ctx, tx, userID) + if err != nil { + return err + } + msgCount, err := e.messagesThisMonthForTx(ctx, tx, userID) + if err != nil { + return err + } + if msgCount+units > lim.MaxMessagesMonth { + return &LimitExceededError{ + Resource: "messages_month", + Limit: lim.MaxMessagesMonth, + Current: msgCount, + Limits: lim, + } + } + if lim.MaxMessagesDay != nil { + dayCount, err := e.messagesTodayForTx(ctx, tx, userID) + if err != nil { + return err + } + if dayCount+units > *lim.MaxMessagesDay { + return &LimitExceededError{ + Resource: "messages_day", + Limit: *lim.MaxMessagesDay, + Current: dayCount, + Limits: lim, + } + } + } + return nil +} + +// limitsForTx resolves the caps without a pool read while the reservation +// holds its advisory lock: the cache first (the accept pre-check has normally +// just filled it), then the transaction, and the pooled getter only for an +// enforcer whose store has no tx form. +func (e *DBEnforcer) limitsForTx(ctx context.Context, tx pgx.Tx, userID string) (Limits, error) { + if cached, _, ok := e.cacheGet(userID); ok { + return cached, nil + } + if r, ok := e.store.(limitsTxReader); ok { + row, found, err := r.GetTx(ctx, tx, userID) + if err != nil { + return Limits{}, err + } + return e.resolveLimits(row, found), nil + } + return e.Get(ctx, userID) +} + +func (e *DBEnforcer) messagesThisMonthForTx(ctx context.Context, tx pgx.Tx, userID string) (int, error) { + if c, ok := e.counter.(counterTxReader); ok { + return c.MessagesThisMonthTx(ctx, tx, userID) + } + return e.counter.MessagesThisMonth(ctx, userID) +} + +func (e *DBEnforcer) messagesTodayForTx(ctx context.Context, tx pgx.Tx, userID string) (int, error) { + if c, ok := e.counter.(counterTxReader); ok { + return c.MessagesTodayTx(ctx, tx, userID) + } + return e.counter.MessagesToday(ctx, userID) +} + // CheckInboundMessage enforces ONLY the storage stock cap for an inbound // delivery. Message-flow caps are outbound-only (usage-based pricing v1): // a recipient's exhausted send allowance must never bounce a stranger's diff --git a/internal/limits/store.go b/internal/limits/store.go index f5950d0f7..4d97924ea 100644 --- a/internal/limits/store.go +++ b/internal/limits/store.go @@ -25,8 +25,25 @@ func NewStore(pool *pgxpool.Pool) *Store { // error" is important: a transient DB hiccup must fail closed, while a // genuinely-missing row is the normal case for a fresh user. func (s *Store) Get(ctx context.Context, userID string) (Limits, bool, error) { + return s.get(ctx, s.pool, userID) +} + +// GetTx is Get on a caller-owned transaction. The accept-time send +// reservation reads the cap on the accept transaction's own connection so it +// never needs a second pool connection while it holds the per-user advisory +// lock. +func (s *Store) GetTx(ctx context.Context, tx pgx.Tx, userID string) (Limits, bool, error) { + return s.get(ctx, tx, userID) +} + +// rowQuerier is the subset of *pgxpool.Pool and pgx.Tx get needs. +type rowQuerier interface { + QueryRow(ctx context.Context, sql string, args ...any) pgx.Row +} + +func (s *Store) get(ctx context.Context, q rowQuerier, userID string) (Limits, bool, error) { l := Limits{} - err := s.pool.QueryRow(ctx, + err := q.QueryRow(ctx, `SELECT plan_code, max_agents, max_domains, max_messages_month, max_messages_day, max_storage_bytes, upgrade_url, outbound_footer_enabled FROM account_limits WHERE user_id = $1`, userID, ).Scan(&l.PlanCode, &l.MaxAgents, &l.MaxDomains, &l.MaxMessagesMonth, &l.MaxMessagesDay, &l.MaxStorageBytes, &l.UpgradeURL, &l.OutboundFooterEnabled) diff --git a/internal/usage/pending_quota_test.go b/internal/usage/pending_quota_test.go new file mode 100644 index 000000000..9bec6ab9e --- /dev/null +++ b/internal/usage/pending_quota_test.go @@ -0,0 +1,111 @@ +package usage_test + +import ( + "context" + "testing" + "time" + + "github.com/jackc/pgx/v5" + + "github.com/tokencanopy/e2a/internal/identity" + "github.com/tokencanopy/e2a/internal/testutil" + "github.com/tokencanopy/e2a/internal/usage" +) + +// Accepted-but-unmetered outbound sends must count against the flow caps: the +// terminal metering write is the only thing that lands in usage_summaries, so +// without counting the accepted rows the cap reads the same pre-increment +// total for every send already in flight. This pins the counter's new +// accept-time reservation semantics and the rows it must NOT count. +func TestMessagesCountsAcceptedButUnmeteredSends(t *testing.T) { + if testing.Short() { + t.Skip("skipping DB-backed usage test under -short") + } + pool := testutil.TestDB(t) + store := usage.NewStore(pool) + idStore := identity.NewStore(pool) + ctx := context.Background() + + user, err := idStore.CreateOrGetUser(ctx, "pending-quota@example.com", "Pending", "google-pending-quota") + if err != nil { + t.Fatalf("CreateOrGetUser: %v", err) + } + if _, err := idStore.ClaimOrCreateDomain(ctx, "pending.example.com", user.ID); err != nil { + t.Fatalf("ClaimOrCreateDomain: %v", err) + } + agent, err := idStore.CreateAgent(ctx, "bot@pending.example.com", "pending.example.com", "", "", "", user.ID) + if err != nil { + t.Fatalf("CreateAgent: %v", err) + } + + // Two accepted immediate sends: one with two recipients, one with one. + accept := func(to []string, schedule time.Time) { + t.Helper() + if err := idStore.WithTx(ctx, func(tx pgx.Tx) error { + msg, err := idStore.CreateOutboundMessageTx(ctx, tx, agent.ID, to, nil, nil, + "pending quota", "send", "smtp", "", "", []byte("raw"), "accepted", "", "") + if err != nil { + return err + } + if !schedule.IsZero() { + return idStore.StampScheduledAtTx(ctx, tx, msg.ID, schedule) + } + return nil + }); err != nil { + t.Fatalf("accept send: %v", err) + } + } + + accept([]string{"alice@example.com", "bob@example.com"}, time.Time{}) + accept([]string{"carol@example.com"}, time.Time{}) + + got, err := store.MessagesThisMonth(ctx, user.ID) + if err != nil { + t.Fatalf("MessagesThisMonth: %v", err) + } + if got != 3 { + t.Errorf("MessagesThisMonth = %d, want 3 (2-recipient + 1-recipient accepted sends)", got) + } + got, err = store.MessagesToday(ctx, user.ID) + if err != nil { + t.Fatalf("MessagesToday: %v", err) + } + if got != 3 { + t.Errorf("MessagesToday = %d, want 3", got) + } + + // A scheduled send's flow cap is judged against the month it fires in, and + // a review hold is re-checked when released: neither is an accept-time + // reservation. + accept([]string{"dave@example.com", "erin@example.com"}, time.Now().UTC().Add(time.Hour)) + if _, err := idStore.CreatePendingOutboundMessage(ctx, agent.ID, []string{"frank@example.com"}, nil, nil, + "held", "body", "", nil, "send", "", "", "", 3600); err != nil { + t.Fatalf("CreatePendingOutboundMessage: %v", err) + } + got, err = store.MessagesThisMonth(ctx, user.ID) + if err != nil { + t.Fatalf("MessagesThisMonth after scheduled + held: %v", err) + } + if got != 3 { + t.Errorf("MessagesThisMonth = %d, want 3 (scheduled and held sends are not reservations)", got) + } + + // Terminal send: the metering write lands in usage_summaries and the row + // leaves the accepted set, so the total must not double count. + if _, err := pool.Exec(ctx, + `UPDATE messages SET delivery_status = 'sent' WHERE agent_id = $1 AND to_recipients = ARRAY['alice@example.com','bob@example.com']`, + agent.ID, + ); err != nil { + t.Fatalf("mark sent: %v", err) + } + if err := store.IncrementUsageSummary(ctx, user.ID, usage.CurrentDate(), "outbound", 2); err != nil { + t.Fatalf("IncrementUsageSummary: %v", err) + } + got, err = store.MessagesThisMonth(ctx, user.ID) + if err != nil { + t.Fatalf("MessagesThisMonth after terminal: %v", err) + } + if got != 3 { + t.Errorf("MessagesThisMonth = %d, want 3 (metered 2 + still-accepted 1, no double count)", got) + } +} diff --git a/internal/usage/store.go b/internal/usage/store.go index 24966d4aa..e9af6f468 100644 --- a/internal/usage/store.go +++ b/internal/usage/store.go @@ -52,6 +52,13 @@ type Store struct { pool *pgxpool.Pool } +// rowQuerier is the subset of *pgxpool.Pool and pgx.Tx the counters need, so +// the same read serves both the pooling callers and the accept-time +// reservation running on the accept transaction. +type rowQuerier interface { + QueryRow(ctx context.Context, sql string, args ...any) pgx.Row +} + func NewStore(pool *pgxpool.Pool) *Store { return &Store{pool: pool} } @@ -207,24 +214,49 @@ func (s *Store) CountDomainsByUser(ctx context.Context, userID string) (int, err } // MessagesThisMonth returns the user's OUTBOUND recipient-delivery count -// for the current UTC calendar month, summed from usage_summaries. -// Inbound mail is recorded (inbound_count feeds analytics and dashboards) -// but is free and unmetered — it does not consume the monthly allowance. -// Returns 0 with no error if the user has no rows yet. The reference is -// time.Now().UTC() so server clocks crossing midnight UTC roll the -// counter consistently with the daily bucket_date written by +// for the current UTC calendar month: the terminally-metered total from +// usage_summaries plus the units of sends that are durably accepted but not +// yet terminal. Inbound mail is recorded (inbound_count feeds analytics and +// dashboards) but is free and unmetered — it does not consume the monthly +// allowance. Returns 0 with no error if the user has no usage yet. The +// reference is time.Now().UTC() so server clocks crossing midnight UTC roll +// the counter consistently with the daily bucket_date written by // IncrementUsageSummary. +// +// The accepted-but-unmetered units are the accept-time quota reservations: +// the metering write happens only at a send's terminal outcome, so without +// counting the accepted rows every send already in flight would be invisible +// to the cap and a burst could all pass against the same pre-increment total. func (s *Store) MessagesThisMonth(ctx context.Context, userID string) (int, error) { + return s.messagesThisMonth(ctx, s.pool, userID) +} + +// MessagesThisMonthTx is MessagesThisMonth on a caller-owned transaction. The +// accept-time reservation reads under the per-user advisory lock, and doing +// that read on the accept transaction's own connection keeps it from needing a +// second pool connection while the lock is held (a saturated pool would +// otherwise deadlock the accept path). +func (s *Store) MessagesThisMonthTx(ctx context.Context, tx pgx.Tx, userID string) (int, error) { + return s.messagesThisMonth(ctx, tx, userID) +} + +func (s *Store) messagesThisMonth(ctx context.Context, q rowQuerier, userID string) (int, error) { now := time.Now().UTC() - monthStart := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, time.UTC).Format("2006-01-02") + monthStart := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, time.UTC) var count int - err := s.pool.QueryRow(ctx, + if err := q.QueryRow(ctx, `SELECT COALESCE(SUM(outbound_count), 0) FROM usage_summaries WHERE user_id = $1 AND bucket_date >= $2`, - userID, monthStart, - ).Scan(&count) - return count, err + userID, monthStart.Format("2006-01-02"), + ).Scan(&count); err != nil { + return 0, err + } + pending, err := s.pendingOutboundUnits(ctx, q, userID, monthStart, monthStart.AddDate(0, 1, 0)) + if err != nil { + return 0, err + } + return count + pending, nil } // MessagesToday returns the user's OUTBOUND recipient-delivery count for @@ -232,16 +264,62 @@ func (s *Store) MessagesThisMonth(ctx context.Context, userID string) (int, erro // no error if the user has no row for today. Shares CurrentDate()'s UTC // bucketing with IncrementUsageSummary so the day rolls consistently. func (s *Store) MessagesToday(ctx context.Context, userID string) (int, error) { + return s.messagesToday(ctx, s.pool, userID) +} + +// MessagesTodayTx is MessagesToday on a caller-owned transaction; see +// MessagesThisMonthTx for why the reservation reads on the accept tx. +func (s *Store) MessagesTodayTx(ctx context.Context, tx pgx.Tx, userID string) (int, error) { + return s.messagesToday(ctx, tx, userID) +} + +func (s *Store) messagesToday(ctx context.Context, q rowQuerier, userID string) (int, error) { var count int - err := s.pool.QueryRow(ctx, + err := q.QueryRow(ctx, `SELECT outbound_count FROM usage_summaries WHERE user_id = $1 AND bucket_date = $2`, userID, CurrentDate(), ).Scan(&count) - if errors.Is(err, pgx.ErrNoRows) { - return 0, nil + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return 0, err } - return count, err + now := time.Now().UTC() + dayStart := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) + pending, err := s.pendingOutboundUnits(ctx, q, userID, dayStart, dayStart.AddDate(0, 0, 1)) + if err != nil { + return 0, err + } + return count + pending, nil +} + +// pendingOutboundUnits returns the recipient-delivery units of outbound +// messages created in [from, until) that are durably accepted but have not +// reached a terminal delivery status. These rows are the accept-time quota +// reservations that make MessagesThisMonth / MessagesToday see a concurrent +// burst. Scheduled sends are excluded — their quota is judged against the +// target month by the fire-time gate, not the month they were accepted in — +// and review holds are excluded because they are re-checked when released. +// Units are the deduplicated to ∪ cc ∪ bcc set, matching +// identity.UniqueRecipientCount and the units IncrementUsageSummary writes. +func (s *Store) pendingOutboundUnits(ctx context.Context, q rowQuerier, userID string, from, until time.Time) (int, error) { + var units int + err := q.QueryRow(ctx, ` + SELECT COALESCE(SUM(( + SELECT count(DISTINCT lower(trim(addr))) + FROM unnest(COALESCE(m.to_recipients, '{}') || COALESCE(m.cc, '{}') || COALESCE(m.bcc, '{}')) AS addr + WHERE trim(addr) <> '' + )), 0) + FROM messages m + JOIN agent_identities a ON a.id = m.agent_id + WHERE a.user_id = $1 + AND m.direction = 'outbound' + AND m.delivery_status IN ('accepted', 'sending', 'queued') + AND m.status IS DISTINCT FROM 'pending_review' + AND m.scheduled_at IS NULL + AND m.created_at >= $2 + AND m.created_at < $3`, + userID, from, until).Scan(&units) + return units, err } // GetStorageBytes returns the user's current materialized storage bytes