From 8c33dc27e4385c640e34fb1a51bb514127e45d9f Mon Sep 17 00:00:00 2001 From: tiennm99 Date: Thu, 11 Jun 2026 21:12:58 +0700 Subject: [PATCH] feat(gold): add compare-and-swap portfolio updates with multi-backend support Implement atomic UpdatePortfolio with retry pattern for concurrent-safe updates. Add CAS support to DynamoDB, Firestore, Memory, and prefix stores. --- cmd/server/main_test.go | 2 +- internal/modules/gold/handlers.go | 68 ++++++++++++++----------- internal/modules/gold/portfolio.go | 50 ++++++++++++++++++ internal/modules/gold/portfolio_test.go | 40 +++++++++++++++ internal/storage/dynamodb_kv.go | 34 +++++++++++++ internal/storage/firestore_kv.go | 50 ++++++++++++++++++ internal/storage/kv_store.go | 10 ++++ internal/storage/memory_kv.go | 17 +++++++ internal/storage/prefix.go | 8 +++ 9 files changed, 249 insertions(+), 30 deletions(-) diff --git a/cmd/server/main_test.go b/cmd/server/main_test.go index 446ab62..b77e9ed 100644 --- a/cmd/server/main_test.go +++ b/cmd/server/main_test.go @@ -16,7 +16,7 @@ func TestFactoriesIncludesGold(t *testing.T) { if err != nil { t.Fatalf("Build gold: %v", err) } - for _, name := range []string{"gold_topup", "gold_buy", "gold_sell", "gold_stats"} { + for _, name := range []string{"gold_price", "gold_topup", "gold_buy", "gold_sell", "gold_stats"} { if _, ok := reg.AllCommands[name]; !ok { t.Fatalf("missing command %s", name) } diff --git a/internal/modules/gold/handlers.go b/internal/modules/gold/handlers.go index d5d524d..8738348 100644 --- a/internal/modules/gold/handlers.go +++ b/internal/modules/gold/handlers.go @@ -2,6 +2,7 @@ package gold import ( "context" + "errors" "strconv" "strings" @@ -12,6 +13,11 @@ import ( "github.com/tiennm99/miti99bot/internal/modules/util/chathelper" ) +var ( + errInsufficientVND = errors.New("gold: insufficient VND") + errInsufficientGold = errors.New("gold: insufficient gold") +) + func (s *state) handlePrice(ctx context.Context, b *bot.Bot, update *models.Update) error { args := argsAfterCommand(update.Message.Text) if len(args) != 0 { @@ -46,14 +52,12 @@ func (s *state) handleTopup(ctx context.Context, b *bot.Bot, update *models.Upda } defer s.locks.Acquire(strconv.FormatInt(userID, 10))() - p, err := LoadPortfolio(ctx, s.kv, userID, s.now().UnixMilli()) + p, err := UpdatePortfolio(ctx, s.kv, userID, s.now().UnixMilli(), func(p *Portfolio) error { + p.AddVND(amount) + p.Meta.Invested += amount + return nil + }) if err != nil { - log.Error("gold_load_portfolio", "user", userID, "err", err) - return chathelper.Reply(ctx, b, update.Message, "Could not load gold portfolio. Try again later.") - } - p.AddVND(amount) - p.Meta.Invested += amount - if err := SavePortfolio(ctx, s.kv, userID, p); err != nil { log.Error("gold_save_portfolio", "user", userID, "err", err) return chathelper.Reply(ctx, b, update.Message, "Could not save gold portfolio. Try again later.") } @@ -85,18 +89,21 @@ func (s *state) handleBuy(ctx context.Context, b *bot.Bot, update *models.Update } defer s.locks.Acquire(strconv.FormatInt(userID, 10))() - p, err := LoadPortfolio(ctx, s.kv, userID, s.now().UnixMilli()) - if err != nil { - log.Error("gold_load_portfolio", "user", userID, "err", err) - return chathelper.Reply(ctx, b, update.Message, "Could not load gold portfolio. Try again later.") - } - ok, balance := p.DeductVND(cost) - if !ok { + var insufficientBalance *float64 + p, err := UpdatePortfolio(ctx, s.kv, userID, s.now().UnixMilli(), func(p *Portfolio) error { + ok, balance := p.DeductVND(cost) + if !ok { + insufficientBalance = &balance + return errInsufficientVND + } + p.AddLuong(qty) + return nil + }) + if errors.Is(err, errInsufficientVND) && insufficientBalance != nil { return chathelper.Reply(ctx, b, update.Message, - "Insufficient VND. Need "+FormatVND(cost)+", have "+FormatVND(balance)+".") + "Insufficient VND. Need "+FormatVND(cost)+", have "+FormatVND(*insufficientBalance)+".") } - p.AddLuong(qty) - if err := SavePortfolio(ctx, s.kv, userID, p); err != nil { + if err != nil { log.Error("gold_save_portfolio", "user", userID, "err", err) return chathelper.Reply(ctx, b, update.Message, "Could not save gold portfolio. Try again later.") } @@ -125,22 +132,25 @@ func (s *state) handleSell(ctx context.Context, b *bot.Bot, update *models.Updat } defer s.locks.Acquire(strconv.FormatInt(userID, 10))() - p, err := LoadPortfolio(ctx, s.kv, userID, s.now().UnixMilli()) - if err != nil { - log.Error("gold_load_portfolio", "user", userID, "err", err) - return chathelper.Reply(ctx, b, update.Message, "Could not load gold portfolio. Try again later.") - } - ok, held := p.DeductLuong(qty) - if !ok { - return chathelper.Reply(ctx, b, update.Message, - "Insufficient gold. You have: "+FormatLuong(held)+" luong") - } revenue := qty * price if !isSafeVND(revenue) { return chathelper.Reply(ctx, b, update.Message, "Trade value is too large.") } - p.AddVND(revenue) - if err := SavePortfolio(ctx, s.kv, userID, p); err != nil { + var insufficientHeld *float64 + p, err := UpdatePortfolio(ctx, s.kv, userID, s.now().UnixMilli(), func(p *Portfolio) error { + ok, held := p.DeductLuong(qty) + if !ok { + insufficientHeld = &held + return errInsufficientGold + } + p.AddVND(revenue) + return nil + }) + if errors.Is(err, errInsufficientGold) && insufficientHeld != nil { + return chathelper.Reply(ctx, b, update.Message, + "Insufficient gold. You have: "+FormatLuong(*insufficientHeld)+" luong") + } + if err != nil { log.Error("gold_save_portfolio", "user", userID, "err", err) return chathelper.Reply(ctx, b, update.Message, "Could not save gold portfolio. Try again later.") } diff --git a/internal/modules/gold/portfolio.go b/internal/modules/gold/portfolio.go index 4f0b0c0..8d3087d 100644 --- a/internal/modules/gold/portfolio.go +++ b/internal/modules/gold/portfolio.go @@ -2,6 +2,7 @@ package gold import ( "context" + "encoding/json" "errors" "fmt" "math" @@ -11,6 +12,7 @@ import ( ) const goldDustEpsilon = 1e-9 +const portfolioUpdateAttempts = 5 type Portfolio struct { VND float64 `json:"vnd"` @@ -56,6 +58,54 @@ func SavePortfolio(ctx context.Context, kv storage.KVStore, userID int64, p Port return nil } +func UpdatePortfolio(ctx context.Context, kv storage.KVStore, userID int64, now int64, mutate func(*Portfolio) error) (Portfolio, error) { + cas, ok := kv.(storage.CompareAndSwapStore) + if !ok { + return Portfolio{}, fmt.Errorf("gold: storage does not support conditional portfolio updates") + } + key := portfolioKey(userID) + for attempt := 0; attempt < portfolioUpdateAttempts; attempt++ { + p, expected, err := loadPortfolioForUpdate(ctx, kv, key, now) + if err != nil { + return Portfolio{}, fmt.Errorf("gold: load portfolio %d: %w", userID, err) + } + if err := mutate(&p); err != nil { + return p, err + } + p.normalize() + next, err := json.Marshal(p) + if err != nil { + return Portfolio{}, fmt.Errorf("gold: save portfolio %d: json encode: %w", userID, err) + } + if err := cas.CompareAndSwap(ctx, key, expected, next); err == nil { + return p, nil + } else if !errors.Is(err, storage.ErrConflict) { + return Portfolio{}, fmt.Errorf("gold: save portfolio %d: %w", userID, err) + } + } + return Portfolio{}, fmt.Errorf("gold: save portfolio %d: %w", userID, storage.ErrConflict) +} + +func loadPortfolioForUpdate(ctx context.Context, kv storage.KVStore, key string, now int64) (Portfolio, []byte, error) { + raw, err := kv.Get(ctx, key) + switch { + case err == nil: + var p Portfolio + if err := json.Unmarshal(raw, &p); err != nil { + return Portfolio{}, nil, fmt.Errorf("json decode: %w", err) + } + p.normalize() + if p.Meta.CreatedAt == 0 { + p.Meta.CreatedAt = now + } + return p, raw, nil + case errors.Is(err, storage.ErrNotFound): + return NewPortfolio(now), nil, nil + default: + return Portfolio{}, nil, err + } +} + func (p *Portfolio) AddVND(amount float64) { p.VND += amount p.normalize() diff --git a/internal/modules/gold/portfolio_test.go b/internal/modules/gold/portfolio_test.go index cc0eb75..f1a8875 100644 --- a/internal/modules/gold/portfolio_test.go +++ b/internal/modules/gold/portfolio_test.go @@ -68,6 +68,46 @@ func TestNormalizeAmountSpecialValues(t *testing.T) { } } +type conflictOnceStore struct { + storage.KVStore + conflicted bool +} + +func (s *conflictOnceStore) CompareAndSwap(ctx context.Context, key string, expected []byte, val []byte) error { + if !s.conflicted { + s.conflicted = true + competing := NewPortfolio(1) + competing.AddVND(10) + if err := s.KVStore.PutJSON(ctx, key, competing); err != nil { + return err + } + return storage.ErrConflict + } + return s.KVStore.(storage.CompareAndSwapStore).CompareAndSwap(ctx, key, expected, val) +} + +func TestUpdatePortfolioRetriesAfterWriteConflict(t *testing.T) { + ctx := context.Background() + kv := &conflictOnceStore{KVStore: storage.NewMemoryKVStore()} + got, err := UpdatePortfolio(ctx, kv, 7, 1, func(p *Portfolio) error { + p.AddVND(5) + return nil + }) + if err != nil { + t.Fatalf("UpdatePortfolio: %v", err) + } + if got.VND != 15 { + t.Fatalf("updated portfolio VND = %v, want 15", got.VND) + } + loaded, err := LoadPortfolio(ctx, kv, 7, 1) + if err != nil { + t.Fatalf("LoadPortfolio: %v", err) + } + if loaded.VND != 15 { + t.Fatalf("stored portfolio VND = %v, want 15", loaded.VND) + } +} + func TestTradingAndGoldPortfolioKeysDoNotCollide(t *testing.T) { ctx := context.Background() provider := storage.NewMemoryProvider() diff --git a/internal/storage/dynamodb_kv.go b/internal/storage/dynamodb_kv.go index 8eacb91..49aa7a4 100644 --- a/internal/storage/dynamodb_kv.go +++ b/internal/storage/dynamodb_kv.go @@ -3,6 +3,7 @@ package storage import ( "context" "encoding/json" + "errors" "fmt" "strconv" "time" @@ -106,6 +107,39 @@ func (s *DynamoDBKVStore) Put(ctx context.Context, key string, val []byte) error return nil } +func (s *DynamoDBKVStore) CompareAndSwap(ctx context.Context, key string, expected []byte, val []byte) error { + if err := validateKey(key); err != nil { + return err + } + input := &dynamodb.PutItemInput{ + TableName: aws.String(s.table), + Item: map[string]types.AttributeValue{ + dynamoPKAttr: &types.AttributeValueMemberS{Value: s.moduleName}, + dynamoSKAttr: &types.AttributeValueMemberS{Value: key}, + dynamoValueAttr: &types.AttributeValueMemberS{Value: string(val)}, + dynamoUpdatedAtAttr: &types.AttributeValueMemberN{Value: strconv.FormatInt(time.Now().UTC().UnixNano(), 10)}, + }, + } + if expected == nil { + input.ConditionExpression = aws.String("attribute_not_exists(pk) AND attribute_not_exists(sk)") + } else { + input.ConditionExpression = aws.String("#v = :expected") + input.ExpressionAttributeNames = map[string]string{"#v": dynamoValueAttr} + input.ExpressionAttributeValues = map[string]types.AttributeValue{ + ":expected": &types.AttributeValueMemberS{Value: string(expected)}, + } + } + _, err := s.client.PutItem(ctx, input) + if err == nil { + return nil + } + var conflict *types.ConditionalCheckFailedException + if errors.As(err, &conflict) { + return ErrConflict + } + return fmt.Errorf("dynamodb compare-and-swap %s/%s: %w", s.moduleName, key, err) +} + // PutJSON marshals val and writes the bytes at key. func (s *DynamoDBKVStore) PutJSON(ctx context.Context, key string, val any) error { raw, err := json.Marshal(val) diff --git a/internal/storage/firestore_kv.go b/internal/storage/firestore_kv.go index c592583..fb56eea 100644 --- a/internal/storage/firestore_kv.go +++ b/internal/storage/firestore_kv.go @@ -1,6 +1,7 @@ package storage import ( + "bytes" "context" "encoding/json" "errors" @@ -138,6 +139,55 @@ func (s *FirestoreKVStore) Put(ctx context.Context, key string, val []byte) erro return nil } +func (s *FirestoreKVStore) CompareAndSwap(ctx context.Context, key string, expected []byte, val []byte) error { + if err := validateKey(key); err != nil { + return err + } + ref := s.doc(key) + err := s.client.RunTransaction(ctx, func(ctx context.Context, tx *firestore.Transaction) error { + snap, err := tx.Get(ref) + if err != nil { + if status.Code(err) == codes.NotFound && expected == nil { + return tx.Set(ref, map[string]any{ + firestoreValueField: val, + firestoreUpdatedAtField: time.Now().UTC(), + }) + } + if status.Code(err) == codes.NotFound { + return ErrConflict + } + return err + } + raw, err := snap.DataAt(firestoreValueField) + if err != nil { + return fmt.Errorf("missing %q field: %w", firestoreValueField, err) + } + var current []byte + switch v := raw.(type) { + case []byte: + current = v + case string: + current = []byte(v) + default: + return fmt.Errorf("unexpected value type %T", raw) + } + if expected == nil || !bytes.Equal(current, expected) { + return ErrConflict + } + return tx.Set(ref, map[string]any{ + firestoreValueField: val, + firestoreUpdatedAtField: time.Now().UTC(), + }) + }) + if err != nil { + if errors.Is(err, ErrConflict) { + return ErrConflict + } + return fmt.Errorf("firestore compare-and-swap %s/%s: %w", s.collection, key, err) + } + return nil +} + // PutJSON marshals val and writes the bytes at key. func (s *FirestoreKVStore) PutJSON(ctx context.Context, key string, val any) error { raw, err := json.Marshal(val) diff --git a/internal/storage/kv_store.go b/internal/storage/kv_store.go index f728b57..c20ee3f 100644 --- a/internal/storage/kv_store.go +++ b/internal/storage/kv_store.go @@ -8,6 +8,9 @@ import ( // ErrNotFound is returned by KVStore implementations when a key has no value. var ErrNotFound = errors.New("storage: key not found") +// ErrConflict is returned by conditional writes when the stored value changed. +var ErrConflict = errors.New("storage: write conflict") + // KVStore is the per-module key-value contract. Implementations must be safe // for concurrent use and must return ErrNotFound for missing keys. type KVStore interface { @@ -18,3 +21,10 @@ type KVStore interface { Delete(ctx context.Context, key string) error List(ctx context.Context, prefix string) ([]string, error) } + +// CompareAndSwapStore is implemented by stores that can conditionally replace +// a value only when it still equals the bytes read by the caller. A nil expected +// value means the key must not exist. +type CompareAndSwapStore interface { + CompareAndSwap(ctx context.Context, key string, expected []byte, val []byte) error +} diff --git a/internal/storage/memory_kv.go b/internal/storage/memory_kv.go index 2bd8b5c..d64f5ac 100644 --- a/internal/storage/memory_kv.go +++ b/internal/storage/memory_kv.go @@ -50,6 +50,23 @@ func (s *MemoryKVStore) Put(_ context.Context, key string, val []byte) error { return nil } +func (s *MemoryKVStore) CompareAndSwap(_ context.Context, key string, expected []byte, val []byte) error { + s.mu.Lock() + defer s.mu.Unlock() + current, ok := s.m[key] + if expected == nil { + if ok { + return ErrConflict + } + } else if !ok || !bytes.Equal(current, expected) { + return ErrConflict + } + stored := make([]byte, len(val)) + copy(stored, val) + s.m[key] = stored + return nil +} + func (s *MemoryKVStore) PutJSON(ctx context.Context, key string, val any) error { raw, err := json.Marshal(val) if err != nil { diff --git a/internal/storage/prefix.go b/internal/storage/prefix.go index 86967ea..56ca00f 100644 --- a/internal/storage/prefix.go +++ b/internal/storage/prefix.go @@ -34,6 +34,14 @@ func (p *prefixedStore) Put(ctx context.Context, key string, val []byte) error { return p.inner.Put(ctx, p.k(key), val) } +func (p *prefixedStore) CompareAndSwap(ctx context.Context, key string, expected []byte, val []byte) error { + cas, ok := p.inner.(CompareAndSwapStore) + if !ok { + return ErrConflict + } + return cas.CompareAndSwap(ctx, p.k(key), expected, val) +} + func (p *prefixedStore) PutJSON(ctx context.Context, key string, val any) error { return p.inner.PutJSON(ctx, p.k(key), val) }