Skip to content
Open
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
27 changes: 15 additions & 12 deletions cloud/aws/dynamo/config.go
Original file line number Diff line number Diff line change
@@ -1,11 +1,12 @@
package dynamo

import (
"context"
"fmt"
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/aws/session"
"github.com/aws/aws-sdk-go/service/dynamodb"
"github.com/aws/aws-sdk-go/service/dynamodb/dynamodbattribute"

"github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/feature/dynamodb/attributevalue"
"github.com/aws/aws-sdk-go-v2/service/dynamodb"
"github.com/pkg/errors"
)

Expand All @@ -15,19 +16,21 @@ const keyColumnName = "key"
type Config struct{}

func (c Config) Query(key string) (interface{}, error) {
sess := session.Must(session.NewSessionWithOptions(session.Options{
SharedConfigState: session.SharedConfigEnable,
}))
connection := dynamodb.New(sess)
cfg, err := config.LoadDefaultConfig(context.TODO())
if err != nil {
return nil, errors.Wrap(err, "could not load aws config")
}
connection := dynamodb.NewFromConfig(cfg)

dbKey, err := dynamodbattribute.MarshalMap(map[string]interface{}{keyColumnName: key})
dbKey, err := attributevalue.MarshalMap(map[string]interface{}{keyColumnName: key})
if err != nil {
return nil, errors.Wrap(err, "could not marshal key for dynamo config table")
}

item, err := connection.GetItem(&dynamodb.GetItemInput{
table := tableName
item, err := connection.GetItem(context.TODO(), &dynamodb.GetItemInput{
Key: dbKey,
TableName: aws.String(tableName),
TableName: &table,
})
if err != nil {
return nil, errors.Wrap(err, "could query dynamo config table")
Expand All @@ -38,7 +41,7 @@ func (c Config) Query(key string) (interface{}, error) {
}

var value interface{}
err = dynamodbattribute.Unmarshal(item.Item["value"], &value)
err = attributevalue.Unmarshal(item.Item["value"], &value)
return value, errors.Wrap(err, "Could not unmarshal dynamo config value")
}

Expand Down
134 changes: 92 additions & 42 deletions cloud/aws/dynamo/dynamo.go
Original file line number Diff line number Diff line change
@@ -1,48 +1,49 @@
package dynamo

import (
"context"
"fmt"
"github.com/Direct-Debit/go-commons/cloud"
"github.com/Direct-Debit/go-commons/concurrency"
"github.com/Direct-Debit/go-commons/stdext"
"github.com/aws/aws-sdk-go/aws"
"github.com/aws/aws-sdk-go/service/dynamodb/dynamodbattribute"
"github.com/pkg/errors"
"github.com/aws/aws-sdk-go-v2/feature/dynamodb/attributevalue"
"math"
"math/rand"
"strings"
"sync"
"time"

"github.com/Direct-Debit/go-commons/errlib"
"github.com/aws/aws-sdk-go/aws/session"
"github.com/aws/aws-sdk-go/service/dynamodb"
"github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/service/dynamodb"
dynamotypes "github.com/aws/aws-sdk-go-v2/service/dynamodb/types"
log "github.com/sirupsen/logrus"
)

var connection *dynamodb.DynamoDB
var connection *dynamodb.Client

type Item map[string]*dynamodb.AttributeValue
type Item map[string]dynamotypes.AttributeValue

func Connect() *dynamodb.DynamoDB {
if connection != (*dynamodb.DynamoDB)(nil) {
func Connect() *dynamodb.Client {
if connection != nil {
return connection
}
sess := session.Must(session.NewSessionWithOptions(session.Options{
SharedConfigState: session.SharedConfigEnable,
}))
connection = dynamodb.New(sess)
cfg, err := config.LoadDefaultConfig(context.TODO())
if err != nil {
panic(err)
}
connection = dynamodb.NewFromConfig(cfg)
return connection
}

func TableExists(tableName *string) bool {
db := Connect()

descOutput, err := db.DescribeTable(&dynamodb.DescribeTableInput{TableName: tableName})
descOutput, err := db.DescribeTable(context.TODO(), &dynamodb.DescribeTableInput{TableName: tableName})
if err != nil {
return false
}
return strings.ToLower(*descOutput.Table.TableStatus) == "active"
return strings.ToLower(string(descOutput.Table.TableStatus)) == "active"
}

func PutItems(items []Item, table *string, delete bool) {
Expand All @@ -54,15 +55,15 @@ func PutItems(items []Item, table *string, delete bool) {
}
submitItems := items[i:max]

writeRequests := make([]*dynamodb.WriteRequest, len(submitItems))
writeRequests := make([]dynamotypes.WriteRequest, len(submitItems))
for idx, item := range submitItems {
if delete {
writeRequests[idx] = &dynamodb.WriteRequest{
DeleteRequest: &dynamodb.DeleteRequest{Key: item},
writeRequests[idx] = dynamotypes.WriteRequest{
DeleteRequest: &dynamotypes.DeleteRequest{Key: item},
}
} else {
writeRequests[idx] = &dynamodb.WriteRequest{
PutRequest: &dynamodb.PutRequest{Item: item},
writeRequests[idx] = dynamotypes.WriteRequest{
PutRequest: &dynamotypes.PutRequest{Item: item},
}
}
}
Expand All @@ -72,21 +73,21 @@ func PutItems(items []Item, table *string, delete bool) {
go func() {
defer wg.Done()
putItemsWithBackoff(
map[string][]*dynamodb.WriteRequest{*table: writeRequests},
map[string][]dynamotypes.WriteRequest{*table: writeRequests},
1)
}()
}
wg.Wait()
}

func putItemsWithBackoff(items map[string][]*dynamodb.WriteRequest, backoff int) {
func putItemsWithBackoff(items map[string][]dynamotypes.WriteRequest, backoff int) {
db := Connect()

if backoff < 1 {
backoff = 1
}

out, err := db.BatchWriteItem(&dynamodb.BatchWriteItemInput{
out, err := db.BatchWriteItem(context.TODO(), &dynamodb.BatchWriteItemInput{
RequestItems: items,
})
errlib.PanicError(err, "Couldn't write batch")
Expand Down Expand Up @@ -126,17 +127,69 @@ func GetItemsLambda(keys []Item, table *string, batchSize int) []Item {
return res
}

// toV1AttributeValue converts a V2 dynamotypes.AttributeValue into the JSON
// shape produced by the old aws-sdk-go v1 dynamodb.AttributeValue struct
// (e.g. {"S": "foo"}, {"N": "123"}, {"M": {...}}), since the cmn-dynamo-get-items
// lambda still expects that wire format and cannot be changed.
func toV1AttributeValue(av dynamotypes.AttributeValue) map[string]interface{} {
switch v := av.(type) {
case *dynamotypes.AttributeValueMemberS:
return map[string]interface{}{"S": v.Value}
case *dynamotypes.AttributeValueMemberN:
return map[string]interface{}{"N": v.Value}
case *dynamotypes.AttributeValueMemberB:
return map[string]interface{}{"B": v.Value}
case *dynamotypes.AttributeValueMemberSS:
return map[string]interface{}{"SS": v.Value}
case *dynamotypes.AttributeValueMemberNS:
return map[string]interface{}{"NS": v.Value}
case *dynamotypes.AttributeValueMemberBS:
return map[string]interface{}{"BS": v.Value}
case *dynamotypes.AttributeValueMemberBOOL:
return map[string]interface{}{"BOOL": v.Value}
case *dynamotypes.AttributeValueMemberNULL:
return map[string]interface{}{"NULL": v.Value}
case *dynamotypes.AttributeValueMemberL:
list := make([]interface{}, len(v.Value))
for i, elem := range v.Value {
list[i] = toV1AttributeValue(elem)
}
return map[string]interface{}{"L": list}
case *dynamotypes.AttributeValueMemberM:
m := make(map[string]interface{}, len(v.Value))
for key, elem := range v.Value {
m[key] = toV1AttributeValue(elem)
}
return map[string]interface{}{"M": m}
default:
return nil
}
}

func toV1Item(item Item) map[string]interface{} {
v1Item := make(map[string]interface{}, len(item))
for key, av := range item {
v1Item[key] = toV1AttributeValue(av)
}
return v1Item
}

func invokeGetItemsLambda(items []Item, table *string, output chan []Item) {
v1Items := make([]map[string]interface{}, len(items))
for i, item := range items {
v1Items[i] = toV1Item(item)
}

input := map[string]interface{}{
"items": items,
"items": v1Items,
"table": *table,
}
result, err := cloud.CallFunc("cmn-dynamo-get-items", input)
errlib.ErrorError(err, "failed to invoke dynamo-get-items lambda")

var resp []Item
for _, item := range result["items"].([]interface{}) {
dynamoItem, err := dynamodbattribute.MarshalMap(item)
dynamoItem, err := attributevalue.MarshalMap(item)
if errlib.ErrorError(err, "failed to marshal attribute values") {
continue
}
Expand All @@ -153,18 +206,18 @@ func GetItems(items []Item, table *string) []Item {
max = len(items)
}

getItems := make([]map[string]*dynamodb.AttributeValue, 0, max-i)
getItems := make([]map[string]dynamotypes.AttributeValue, 0, max-i)
for item := i; item < max; item++ {
getItems = append(getItems, items[item])
}

keysAndAttr := &dynamodb.KeysAndAttributes{
keysAndAttr := dynamotypes.KeysAndAttributes{
Keys: getItems,
}

log.Debugf("Getting %d items from dynamo table %s", len(getItems), *table)
subRes := getItemsWithBackoff(
map[string]*dynamodb.KeysAndAttributes{*table: keysAndAttr}, 1,
map[string]dynamotypes.KeysAndAttributes{*table: keysAndAttr}, 1,
)

res = append(res, subRes[*table]...)
Expand All @@ -183,17 +236,17 @@ func GetItemsConcurrent(workers int, items []Item, table *string) []Item {
}()

keys := stdext.Map(chunk,
func(item Item) map[string]*dynamodb.AttributeValue {
func(item Item) map[string]dynamotypes.AttributeValue {
return item
},
)
keysAndAttr := &dynamodb.KeysAndAttributes{
keysAndAttr := dynamotypes.KeysAndAttributes{
Keys: keys,
}

log.Debugf("Getting %d items from dynamo table %s", len(chunk), *table)
itemsReturned := getItemsWithBackoff(
map[string]*dynamodb.KeysAndAttributes{*table: keysAndAttr}, 1,
map[string]dynamotypes.KeysAndAttributes{*table: keysAndAttr}, 1,
)
ret, success = itemsReturned[*table], true
return
Expand All @@ -202,14 +255,14 @@ func GetItemsConcurrent(workers int, items []Item, table *string) []Item {
return stdext.Flatten(returnedChunks)
}

func getItemsWithBackoff(items map[string]*dynamodb.KeysAndAttributes, backoff int) map[string][]Item {
func getItemsWithBackoff(items map[string]dynamotypes.KeysAndAttributes, backoff int) map[string][]Item {
db := Connect()

if backoff < 1 {
backoff = 1
}

out, err := db.BatchGetItem(&dynamodb.BatchGetItemInput{
out, err := db.BatchGetItem(context.TODO(), &dynamodb.BatchGetItemInput{
RequestItems: items,
})
errlib.PanicError(err, "Couldn't get batch")
Expand Down Expand Up @@ -243,7 +296,7 @@ func QueryAll(initialInput *dynamodb.QueryInput) ([]Item, error) {
queryDone := false
db := Connect()
for !queryDone {
result, err := db.Query(initialInput)
result, err := db.Query(context.TODO(), initialInput)
if err != nil {
return res, err
}
Expand All @@ -264,21 +317,18 @@ type CountOutput struct {
}

func CountAll(initialInput *dynamodb.QueryInput) (CountOutput, error) {
initialInput.Select = aws.String(dynamodb.SelectCount)
initialInput.Select = dynamotypes.SelectCount
res := CountOutput{}

queryDone := false
db := Connect()
for !queryDone {
result, err := db.Query(initialInput)
result, err := db.Query(context.TODO(), initialInput)
if err != nil {
return res, err
}
if result.Count == (*int64)(nil) || result.ScannedCount == (*int64)(nil) {
return CountOutput{}, errors.New("count or scanned count is nil")
}
res.Count += *result.Count
res.ScannedCount += *result.ScannedCount
res.Count += int64(result.Count)
res.ScannedCount += int64(result.ScannedCount)

queryDone = len(result.LastEvaluatedKey) == 0
initialInput.ExclusiveStartKey = result.LastEvaluatedKey
Expand All @@ -293,7 +343,7 @@ func ScanAll(initialInput *dynamodb.ScanInput) ([]Item, error) {
scanDone := false
db := Connect()
for !scanDone {
result, err := db.Scan(initialInput)
result, err := db.Scan(context.TODO(), initialInput)
if err != nil {
return res, err
}
Expand Down
Loading