mirror of
https://github.com/tiennm99/tiennm99bot.git
synced 2026-10-11 03:13:46 +00:00
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:
1 parent
2501deb4f8
commit
8c33dc27e4
9 files changed
+249
-30
No files matched your search
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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.")
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in new issue
Block a user