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.
This commit is contained in:
tiennm99 committed 2026-06-11 21:12:58 +07:00
1 parent 2501deb4f8
commit 8c33dc27e4
9 files changed
+249 -30

No files matched your search

+1 -1
View File
@@ -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)
}
+39 -29
View File
@@ -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.")
}
+50
View File
@@ -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()
+40
View File
@@ -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()
+34
View File
@@ -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)
+50
View File
@@ -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)
+10
View File
@@ -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
}
+17
View File
@@ -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 {
+8
View File
@@ -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)
}