diff --git a/wasm/main.go b/wasm/main.go index ba55520..1307712 100644 --- a/wasm/main.go +++ b/wasm/main.go @@ -4,6 +4,7 @@ package main import ( + "encoding/hex" "fmt" "strconv" "syscall/js" @@ -87,6 +88,30 @@ func safeUint32(v js.Value, index int) (uint32, error) { return uint32(v.Int()), nil } +// safeInt16 safely extracts an int16 from a js.Value, handling undefined values +func safeInt16(v js.Value, index int) (int16, error) { + if v.Type() == js.TypeUndefined { + return 0, fmt.Errorf("argument %d is undefined", index) + } + return int16(v.Int()), nil +} + +// safeUint64 safely extracts a uint64 from a js.Value, handling undefined values +func safeUint64(v js.Value, index int) (uint64, error) { + if v.Type() == js.TypeUndefined { + return 0, fmt.Errorf("argument %d is undefined", index) + } + return uint64(v.Int()), nil +} + +// safeUint16 safely extracts a uint16 from a js.Value, handling undefined values +func safeUint16(v js.Value, index int) (uint16, error) { + if v.Type() == js.TypeUndefined { + return 0, fmt.Errorf("argument %d is undefined", index) + } + return uint16(v.Int()), nil +} + func getClient(args []js.Value) (*client.TxClient, error) { l := len(args) if l < 2 { @@ -135,11 +160,12 @@ func main() { } url := args[0].String() privateKey := args[1].String() - chainId := uint32(args[2].Int()) + chainIdVal := uint32(args[2].Int()) apiKeyIndex := uint8(args[3].Int()) accountIndex := int64(args[4].Int()) httpClient := http.NewClient(url) - _, err := client.CreateClient(httpClient, privateKey, chainId, apiKeyIndex, accountIndex) + chainId = chainIdVal + _, err := client.CreateClient(httpClient, privateKey, chainIdVal, apiKeyIndex, accountIndex) if err != nil { return wrapErr(err) } @@ -239,7 +265,7 @@ func main() { return wrapErr(err) } - marketIndex, err := safeInt(args[0], 0) + marketIndex, err := safeInt16(args[0], 0) if err != nil { return wrapErr(err) } @@ -335,7 +361,10 @@ func main() { return wrapErr(err) } - marketIndex := int16(args[0].Int()) + marketIndex, err := safeInt16(args[0], 0) + if err != nil { + return wrapErr(err) + } orderIndex := int64(args[1].Int()) nonce := int64(args[2].Int()) @@ -383,33 +412,88 @@ func main() { js.Global().Set("SignTransfer", js.FuncOf(func(this js.Value, args []js.Value) interface{} { return recoverPanic(func() js.Value { - if len(args) < 7 { - return js.ValueOf(map[string]interface{}{"error": "SignTransfer expects 7 args: toAccount, usdcAmount, fee, memo, nonce, apiKeyIndex, accountIndex"}) + if len(args) < 10 { + return js.ValueOf(map[string]interface{}{"error": "SignTransfer expects 10 args: toAccountIndex, assetIndex, fromRouteType, toRouteType, amount, usdcFee, memo, nonce, apiKeyIndex, accountIndex"}) + } + // Validate all arguments are defined before accessing + for i := 0; i < 10; i++ { + if args[i].Type() == js.TypeUndefined { + return js.ValueOf(map[string]interface{}{"error": fmt.Sprintf("argument %d is undefined", i)}) + } } c, err := getClient(args) if err != nil { return wrapErr(err) } - toAccount := int64(args[0].Int()) - usdcAmount := int64(args[1].Int()) - fee := int64(args[2].Int()) - memoStr := args[3].String() - nonce := int64(args[4].Int()) + toAccountIndex, err := safeInt(args[0], 0) + if err != nil { + return wrapErr(err) + } + assetIndex, err := safeInt16(args[1], 1) + if err != nil { + return wrapErr(err) + } + fromRouteType, err := safeUint8(args[2], 2) + if err != nil { + return wrapErr(err) + } + toRouteType, err := safeUint8(args[3], 3) + if err != nil { + return wrapErr(err) + } + amount, err := safeInt(args[4], 4) + if err != nil { + return wrapErr(err) + } + usdcFee, err := safeInt(args[5], 5) + if err != nil { + return wrapErr(err) + } + memoStr := args[6].String() + nonce, err := safeInt(args[7], 7) + if err != nil { + return wrapErr(err) + } + // Memo hex encoding support (matching sharedlib implementation) var memoArr [32]byte - bs := []byte(memoStr) - if len(bs) != 32 { - return wrapErr(fmt.Errorf("memo expected to be 32 bytes long")) + if len(memoStr) == 66 { + if memoStr[0:2] == "0x" { + memoStr = memoStr[2:66] + } else { + return wrapErr(fmt.Errorf("memo expected to be 32 bytes or 64 hex encoded or 66 if 0x hex encoded -- long but received %v", len(memoStr))) + } } - for i := 0; i < 32; i++ { - memoArr[i] = bs[i] + + // assume hex encoded here + if len(memoStr) == 64 { + b, err := hex.DecodeString(memoStr) + if err != nil { + return wrapErr(fmt.Errorf("failed to decode hex string. err: %v", err)) + } + if len(b) != 32 { + return wrapErr(fmt.Errorf("decoded hex string must be 32 bytes, got %d", len(b))) + } + for i := 0; i < 32; i++ { + memoArr[i] = b[i] + } + } else if len(memoStr) == 32 { + bs := []byte(memoStr) + for i := 0; i < 32; i++ { + memoArr[i] = bs[i] + } + } else { + return wrapErr(fmt.Errorf("memo expected to be 32 bytes or 64 hex encoded or 66 if 0x hex encoded -- long but received %v", len(memoStr))) } txInfo := &types.TransferTxReq{ - ToAccountIndex: toAccount, - Amount: usdcAmount, - USDCFee: fee, + ToAccountIndex: toAccountIndex, + AssetIndex: assetIndex, + FromRouteType: fromRouteType, + ToRouteType: toRouteType, + Amount: amount, + USDCFee: usdcFee, Memo: memoArr, } ops := new(types.TransactOpts) @@ -424,19 +508,41 @@ func main() { js.Global().Set("SignWithdraw", js.FuncOf(func(this js.Value, args []js.Value) interface{} { return recoverPanic(func() js.Value { - if len(args) < 4 { - return js.ValueOf(map[string]interface{}{"error": "SignWithdraw expects 4 args: usdcAmount, nonce, apiKeyIndex, accountIndex"}) + if len(args) < 6 { + return js.ValueOf(map[string]interface{}{"error": "SignWithdraw expects 6 args: assetIndex, routeType, amount, nonce, apiKeyIndex, accountIndex"}) + } + // Validate all arguments are defined before accessing + for i := 0; i < 6; i++ { + if args[i].Type() == js.TypeUndefined { + return js.ValueOf(map[string]interface{}{"error": fmt.Sprintf("argument %d is undefined", i)}) + } } c, err := getClient(args) if err != nil { return wrapErr(err) } - usdcAmount := uint64(args[0].Int()) - nonce := int64(args[1].Int()) + assetIndex, err := safeInt16(args[0], 0) + if err != nil { + return wrapErr(err) + } + routeType, err := safeUint8(args[1], 1) + if err != nil { + return wrapErr(err) + } + amount, err := safeUint64(args[2], 2) + if err != nil { + return wrapErr(err) + } + nonce, err := safeInt(args[3], 3) + if err != nil { + return wrapErr(err) + } txInfo := &types.WithdrawTxReq{ - Amount: usdcAmount, + AssetIndex: assetIndex, + RouteType: routeType, + Amount: amount, } ops := new(types.TransactOpts) if nonce != -1 { @@ -458,7 +564,10 @@ func main() { return wrapErr(err) } - marketIndex := int16(args[0].Int()) + marketIndex, err := safeInt16(args[0], 0) + if err != nil { + return wrapErr(err) + } fraction := uint16(args[1].Int()) marginMode := uint8(args[2].Int()) nonce := int64(args[3].Int()) @@ -488,7 +597,10 @@ func main() { return wrapErr(err) } - marketIndex := int16(args[0].Int()) + marketIndex, err := safeInt16(args[0], 0) + if err != nil { + return wrapErr(err) + } index := int64(args[1].Int()) baseAmount := int64(args[2].Int()) price := uint32(args[3].Int()) @@ -552,7 +664,10 @@ func main() { operatorFee := int64(args[0].Int()) initialTotalShares := int64(args[1].Int()) - minOperatorShareRate := uint16(args[2].Int()) + minOperatorShareRate, err := safeUint16(args[2], 2) + if err != nil { + return wrapErr(err) + } nonce := int64(args[3].Int()) txInfo := &types.CreatePublicPoolTxReq{ @@ -580,14 +695,20 @@ func main() { return wrapErr(err) } - publicPoolIndex := uint8(args[0].Int()) + publicPoolIndex, err := safeInt(args[0], 0) + if err != nil { + return wrapErr(err) + } status := uint8(args[1].Int()) operatorFee := int64(args[2].Int()) - minOperatorShareRate := uint16(args[3].Int()) + minOperatorShareRate, err := safeUint16(args[3], 3) + if err != nil { + return wrapErr(err) + } nonce := int64(args[4].Int()) txInfo := &types.UpdatePublicPoolTxReq{ - PublicPoolIndex: int64(publicPoolIndex), + PublicPoolIndex: publicPoolIndex, Status: status, OperatorFee: operatorFee, MinOperatorShareRate: minOperatorShareRate, @@ -724,7 +845,10 @@ func main() { return wrapErr(err) } - marketIndex := int16(args[0].Int()) + marketIndex, err := safeInt16(args[0], 0) + if err != nil { + return wrapErr(err) + } usdcAmount := int64(args[1].Int()) direction := uint8(args[2].Int()) nonce := int64(args[3].Int())