Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 10 additions & 5 deletions admin/billing/orb.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,10 +35,11 @@ var ErrCustomerIDRequired = errors.New("customer id is required")
var _ Biller = &Orb{}

type Orb struct {
client *orb.Client
logger *zap.Logger
webhookSecret string
taxProvider string
client *orb.Client
logger *zap.Logger
webhookSecret string
taxProvider string
usageRetryBackoff []time.Duration
}

func NewOrb(logger *zap.Logger, orbKey, webhookSecret, taxProvider string) Biller {
Expand Down Expand Up @@ -672,7 +673,11 @@ func (o *Orb) getAllPlans(ctx context.Context) ([]*Plan, error) {
}

func (o *Orb) pushUsage(ctx context.Context, usage *[]orb.EventIngestParamsEvent) error {
re := retrier.New(retrier.ExponentialBackoff(5, 500*time.Millisecond), retryErrClassifier{})
backoff := o.usageRetryBackoff
if backoff == nil {
backoff = retrier.ExponentialBackoff(5, 500*time.Millisecond)
}
re := retrier.New(backoff, retryErrClassifier{})
err := re.RunCtx(ctx, func(ctx context.Context) error {
resp, err := o.client.Events.Ingest(ctx, orb.EventIngestParams{
Events: orb.F(*usage),
Expand Down
230 changes: 230 additions & 0 deletions admin/billing/orb_usage_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,230 @@
package billing

import (
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"sync"
"testing"
"time"

"github.com/orbcorp/orb-go"
"github.com/orbcorp/orb-go/option"
"github.com/stretchr/testify/require"
"go.uber.org/zap"
)

func TestOrbReportUsageBatchBoundaries(t *testing.T) {
// Orb accepts at most 500 events per ingestion request. These exact boundaries
// guard against empty requests, dropped remainders, and accidental 501-event batches.
tests := []struct {
count int
wantBatches []int
}{
{count: 0, wantBatches: nil},
{count: 1, wantBatches: []int{1}},
{count: 500, wantBatches: []int{500}},
{count: 501, wantBatches: []int{500, 1}},
}
for _, tt := range tests {
t.Run(fmt.Sprintf("%d events", tt.count), func(t *testing.T) {
server := newOrbUsageServer(t, func(_ int, _ orbUsageRequest) (int, string) {
return http.StatusOK, `{"validation_failed":[]}`
})
biller := newTestOrbUsageBiller(server.server.URL)

err := biller.ReportUsage(t.Context(), makeOrbUsageFixtures(tt.count))
require.NoError(t, err)
require.Equal(t, tt.wantBatches, server.batchSizes())
})
}
}

func TestOrbReportUsageSerializesStableLogicalEvent(t *testing.T) {
// Retries and repeated checkpoints must serialize the same logical bucket to
// the same key, regardless of metadata map insertion order.
server := newOrbUsageServer(t, func(_ int, _ orbUsageRequest) (int, string) {
return http.StatusOK, `{"validation_failed":[]}`
})
biller := newTestOrbUsageBiller(server.server.URL)
end := time.Date(2026, time.January, 2, 4, 0, 0, 0, time.UTC)
first := &Usage{
CustomerID: "org-1", MetricName: "api_calls", Value: 42,
ReportingGrain: UsageReportingGranularityHour, EndTime: end,
Metadata: map[string]interface{}{"region": "us-east", "tier": "pro"},
}
second := &Usage{
CustomerID: "org-1", MetricName: "api_calls", Value: 42,
ReportingGrain: UsageReportingGranularityHour, EndTime: end,
Metadata: map[string]interface{}{"tier": "pro", "region": "us-east"},
}

require.NoError(t, biller.ReportUsage(t.Context(), []*Usage{first}))
require.NoError(t, biller.ReportUsage(t.Context(), []*Usage{second}))
requests := server.requestsSnapshot()
require.Len(t, requests, 2)
require.Len(t, requests[0].Events, 1)
require.Equal(t, requests[0].Events[0].IdempotencyKey, requests[1].Events[0].IdempotencyKey)
require.Equal(t, "api_calls_hour", requests[0].Events[0].EventName)
require.Equal(t, "org-1", requests[0].Events[0].ExternalCustomerID)
require.Equal(t, end.Add(-time.Second), requests[0].Events[0].Timestamp)
require.EqualValues(t, 42, requests[0].Events[0].Properties["amount"])
require.Equal(t, "us-east", requests[0].Events[0].Properties["region"])
}

func TestOrbReportUsageCorrectionKeepsLogicalIdempotencyKey(t *testing.T) {
// A corrected value represents the same customer/metric/time bucket. Keeping
// its identity stable prevents a correction retry from becoming an additive event.
server := newOrbUsageServer(t, func(_ int, _ orbUsageRequest) (int, string) {
return http.StatusOK, `{"validation_failed":[]}`
})
biller := newTestOrbUsageBiller(server.server.URL)
fixtures := makeOrbUsageFixtures(1)
require.NoError(t, biller.ReportUsage(t.Context(), fixtures))
fixtures[0].Value = 99
require.NoError(t, biller.ReportUsage(t.Context(), fixtures))

requests := server.requestsSnapshot()
require.Equal(t, requests[0].Events[0].IdempotencyKey, requests[1].Events[0].IdempotencyKey)
require.EqualValues(t, 1, requests[0].Events[0].Properties["amount"])
require.EqualValues(t, 99, requests[1].Events[0].Properties["amount"])
}

func TestOrbReportUsageRetriesOnlyTransientProviderFailures(t *testing.T) {
// The generated Orb client is configured without its own retry loop here so
// the application retry classifier and its exact request count are observable.
tests := []struct {
name string
firstStatus int
wantRequests int
wantErr bool
}{
{name: "500 retries", firstStatus: http.StatusInternalServerError, wantRequests: 2},
{name: "429 retries", firstStatus: http.StatusTooManyRequests, wantRequests: 2},
{name: "400 fails immediately", firstStatus: http.StatusBadRequest, wantRequests: 1, wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
server := newOrbUsageServer(t, func(attempt int, _ orbUsageRequest) (int, string) {
if attempt == 1 {
return tt.firstStatus, orbUsageErrorBody(tt.firstStatus)
}
return http.StatusOK, `{"validation_failed":[]}`
})
biller := newTestOrbUsageBiller(server.server.URL)

err := biller.ReportUsage(t.Context(), makeOrbUsageFixtures(1))
if tt.wantErr {
require.Error(t, err)
} else {
require.NoError(t, err)
}
require.Equal(t, tt.wantRequests, len(server.requestsSnapshot()))
})
}
}

func orbUsageErrorBody(status int) string {
return fmt.Sprintf(`{"status":%d,"title":"provider failure","type":"test_error","validation_errors":[],"detail":"provider failure"}`, status)
}

func TestOrbReportUsageReturnsRecordValidationDetails(t *testing.T) {
// A 200 response may still reject individual records; the error must identify
// their idempotency keys so operators can repair the correct usage checkpoint.
server := newOrbUsageServer(t, func(_ int, request orbUsageRequest) (int, string) {
return http.StatusOK, fmt.Sprintf(`{"validation_failed":[{"idempotency_key":%q,"validation_errors":["amount must be positive"]}]}`, request.Events[0].IdempotencyKey)
})
biller := newTestOrbUsageBiller(server.server.URL)

err := biller.ReportUsage(t.Context(), makeOrbUsageFixtures(1))
require.ErrorContains(t, err, "validation failure for 1 events")
require.ErrorContains(t, err, "amount must be positive")
require.ErrorContains(t, err, "org-0")
require.Len(t, server.requestsSnapshot(), 1)
}

type orbUsageRequest struct {
Events []struct {
EventName string `json:"event_name"`
IdempotencyKey string `json:"idempotency_key"`
ExternalCustomerID string `json:"external_customer_id"`
Timestamp time.Time `json:"timestamp"`
Properties map[string]interface{} `json:"properties"`
} `json:"events"`
}

type orbUsageServer struct {
server *httptest.Server

mu sync.Mutex
requests []orbUsageRequest
}

func newOrbUsageServer(t *testing.T, response func(attempt int, request orbUsageRequest) (int, string)) *orbUsageServer {
t.Helper()
fixture := &orbUsageServer{}
fixture.server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
require.Equal(t, http.MethodPost, r.Method)
require.Equal(t, "/ingest", r.URL.Path)
var request orbUsageRequest
require.NoError(t, json.NewDecoder(r.Body).Decode(&request))
fixture.mu.Lock()
fixture.requests = append(fixture.requests, request)
attempt := len(fixture.requests)
fixture.mu.Unlock()
status, body := response(attempt, request)
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_, _ = w.Write([]byte(body))
}))
t.Cleanup(fixture.server.Close)
return fixture
}

func (s *orbUsageServer) requestsSnapshot() []orbUsageRequest {
s.mu.Lock()
defer s.mu.Unlock()
return append([]orbUsageRequest(nil), s.requests...)
}

func (s *orbUsageServer) batchSizes() []int {
requests := s.requestsSnapshot()
if len(requests) == 0 {
return nil
}
res := make([]int, len(requests))
for i, request := range requests {
res[i] = len(request.Events)
}
return res
}

func newTestOrbUsageBiller(baseURL string) *Orb {
client := orb.NewClient(
option.WithAPIKey("test-key"),
option.WithBaseURL(baseURL),
option.WithMaxRetries(0),
)
return &Orb{
client: client,
logger: zap.NewNop(),
usageRetryBackoff: []time.Duration{0, 0, 0, 0, 0},
}
}

func makeOrbUsageFixtures(count int) []*Usage {
usage := make([]*Usage, count)
for i := range usage {
usage[i] = &Usage{
CustomerID: fmt.Sprintf("org-%d", i),
MetricName: "api_calls",
Value: float64(i + 1),
ReportingGrain: UsageReportingGranularityHour,
StartTime: time.Date(2026, time.January, 2, 3, 0, 0, 0, time.UTC),
EndTime: time.Date(2026, time.January, 2, 4, 0, 0, 0, time.UTC),
Metadata: map[string]interface{}{"region": "us-east"},
}
}
return usage
}
42 changes: 28 additions & 14 deletions admin/billing/orb_webhook.go
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,10 @@ func (o *orbWebhook) handleWebhook(w http.ResponseWriter, r *http.Request) error
r.Body = http.MaxBytesReader(w, r.Body, maxBodyBytes)
payload, err := io.ReadAll(r.Body)
if err != nil {
var maxBytesErr *http.MaxBytesError
if errors.As(err, &maxBytesErr) {
return httputil.Errorf(http.StatusRequestEntityTooLarge, "webhook request body exceeds %d bytes", maxBodyBytes)
}
return httputil.Errorf(http.StatusServiceUnavailable, "error reading request body: %w", err)
}

Expand All @@ -61,12 +65,17 @@ func (o *orbWebhook) handleWebhook(w http.ResponseWriter, r *http.Request) error
return nil
}

// The generic type above is only used to discard events we do not consume.
// Before a recognized event can cause a side effect, verify the signature
// against the exact bytes Orb sent; re-marshaling would change those bytes.
now := time.Now().UTC()
err = o.verifySignature(payload, r.Header, now)
if err != nil {
return httputil.Errorf(http.StatusBadRequest, "error verifying webhook signature: %w", err)
}

// Do not acknowledge work-bearing events until their durable job is queued.
// A queue error becomes a 5xx so Orb retries instead of silently losing work.
switch e.Type {
case "invoice.payment_succeeded":
var ie invoiceEvent
Expand Down Expand Up @@ -154,6 +163,9 @@ func (o *orbWebhook) handleWebhook(w http.ResponseWriter, r *http.Request) error
return nil
}

// Orb delivers webhooks at least once. These handlers deliberately leave replay
// detection to the durable jobs layer; Duplicate means the work is already
// recorded, so it is safe to acknowledge the delivery.
func (o *orbWebhook) handleInvoicePaymentSucceeded(ctx context.Context, ie invoiceEvent) error {
res, err := o.jobs.PaymentSuccess(ctx, ie.OrbInvoice.Customer.ExternalCustomerID, ie.OrbInvoice.ID)
if err != nil {
Expand Down Expand Up @@ -269,20 +281,22 @@ func (o *orbWebhook) verifySignature(payload []byte, headers http.Header, now ti
mac.Write(payload)
expected := mac.Sum(nil)

for _, part := range msgSignature {
parts := strings.Split(part, "=")
if len(parts) != 2 {
continue
}
if parts[0] != "v1" {
continue
}
signature, err := hex.DecodeString(parts[1])
if err != nil {
continue
}
if hmac.Equal(signature, expected) {
return nil
// Orb can include more than one signature during secret rotation. HTTP
// intermediaries may preserve them as repeated fields or combine them with
// commas, so accept either representation when any v1 signature matches.
for _, value := range msgSignature {
for _, part := range strings.Split(value, ",") {
version, encoded, ok := strings.Cut(strings.TrimSpace(part), "=")
if !ok || version != "v1" {
continue
}
signature, err := hex.DecodeString(encoded)
if err != nil {
continue
}
if hmac.Equal(signature, expected) {
return nil
}
}
}

Expand Down
Loading
Loading