diff --git a/cashu/nuts/nut13/nut13.go b/cashu/nuts/nut13/nut13.go index 5020c81..588f4cd 100644 --- a/cashu/nuts/nut13/nut13.go +++ b/cashu/nuts/nut13/nut13.go @@ -1,26 +1,49 @@ package nut13 import ( - "crypto/sha256" - "encoding/binary" + "encoding/base64" "encoding/hex" + "errors" + "fmt" + "math/big" + "regexp" "github.com/btcsuite/btcd/btcutil/hdkeychain" "github.com/decred/dcrd/dcrec/secp256k1/v4" ) +var ( + ErrCollidingKeysetId = errors.New("error: colliding keyset detected") +) + +func keysetIdToBigInt(id string) (*big.Int, error) { + hexPattern := regexp.MustCompile("^[0-9a-fA-F]+$") + + var result *big.Int + modulus := big.NewInt(2147483647) // 2^31 - 1 + + if hexPattern.MatchString(id) { + result = new(big.Int) + result.SetString(id, 16) + } else { + decoded, err := base64.StdEncoding.DecodeString(id) + if err != nil { + return nil, err + } + + hexStr := hex.EncodeToString(decoded) + result = new(big.Int) + result.SetString(hexStr, 16) + } + + return result.Mod(result, modulus), nil +} + func DeriveKeysetPath(master *hdkeychain.ExtendedKey, keysetId string) (*hdkeychain.ExtendedKey, error) { - keysetBytes, err := hex.DecodeString(keysetId) + keysetIdInt, err := keysetIdToBigInt(keysetId) if err != nil { return nil, err } - var keysetIdInt uint64 - if len(keysetBytes) <= 8 { - keysetIdInt = binary.BigEndian.Uint64(keysetBytes) % (1<<31 - 1) - } else { - h := sha256.Sum256(keysetBytes) - keysetIdInt = binary.BigEndian.Uint64(h[:8]) % (1<<31 - 1) - } // m/129372 purpose, err := master.Derive(hdkeychain.HardenedKeyStart + 129372) @@ -35,7 +58,7 @@ func DeriveKeysetPath(master *hdkeychain.ExtendedKey, keysetId string) (*hdkeych } // m/129372'/0'/keyset_k_int' - keysetPath, err := coinType.Derive(hdkeychain.HardenedKeyStart + uint32(keysetIdInt)) + keysetPath, err := coinType.Derive(hdkeychain.HardenedKeyStart + uint32(keysetIdInt.Uint64())) if err != nil { return nil, err } @@ -87,3 +110,29 @@ func DeriveSecret(keysetPath *hdkeychain.ExtendedKey, counter uint32) (string, e return secret, nil } + +func CheckCollidingKeysets(currentKeysetIds []string, newMintKeysetIds []string) error { + for i := range currentKeysetIds { + keysetIdInt, err := keysetIdToBigInt(currentKeysetIds[i]) + if err != nil { + return err + } + + for j := range newMintKeysetIds { + if currentKeysetIds[i] == newMintKeysetIds[j] { + return fmt.Errorf("%w. KeysetId: %+v. New KeysetId: %+v", ErrCollidingKeysetId, currentKeysetIds[i], newMintKeysetIds[j]) + } + + keysetIdIntToCompare, err := keysetIdToBigInt(newMintKeysetIds[j]) + if err != nil { + return err + } + + if keysetIdInt.Cmp(keysetIdIntToCompare) == 0 { + return fmt.Errorf("%w. KeysetId: %+v. New KeysetId: %+v", ErrCollidingKeysetId, currentKeysetIds[i], newMintKeysetIds[j]) + } + } + } + + return nil +} diff --git a/cashu/nuts/nut13/nut13_test.go b/cashu/nuts/nut13/nut13_test.go index 34c6a2f..0a71033 100644 --- a/cashu/nuts/nut13/nut13_test.go +++ b/cashu/nuts/nut13/nut13_test.go @@ -2,6 +2,7 @@ package nut13 import ( "encoding/hex" + "math/big" "testing" "github.com/btcsuite/btcd/btcutil/hdkeychain" @@ -72,3 +73,77 @@ func TestSecretDerivation(t *testing.T) { } } + +func TestDeriveKeysetPath_V2KeyId(t *testing.T) { + mnemonic := "half depart obvious quality work element tank gorilla view sugar picture humble" + v2KeysetId := "01df97b6fb8a572a718d7df7fcbf4387e2d455134ea8004c9c8c51e1b3391f909e" + + seed := bip39.NewSeed(mnemonic, "") + master, err := hdkeychain.NewMaster(seed, &chaincfg.MainNetParams) + if err != nil { + t.Fatal(err) + } + + _, err = DeriveKeysetPath(master, v2KeysetId) + if err != nil { + t.Fatalf("V2 keyset ID derivation failed: %v", err) + } +} + +func TestKeysetIdToBigInt_HexInput(t *testing.T) { + result, err := keysetIdToBigInt("009a1f293253e41e") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if result.Sign() <= 0 { + t.Fatalf("expected positive big.Int, got: %s", result.String()) + } + maxMod := big.NewInt(2147483647) + if result.Cmp(maxMod) > 0 { + t.Fatalf("result should be < 2^31-1, got: %s", result.String()) + } +} + +func TestKeysetIdToBigInt_V2HexInput(t *testing.T) { + v2Id := "01df97b6fb8a572a718d7df7fcbf4387e2d455134ea8004c9c8c51e1b3391f909e" + result, err := keysetIdToBigInt(v2Id) + if err != nil { + t.Fatalf("V2 hex keyset ID failed: %v", err) + } + maxMod := big.NewInt(2147483647) + if result.Cmp(maxMod) > 0 { + t.Fatalf("result should be < 2^31-1, got: %s", result.String()) + } +} + +func TestCheckCollidingKeysets_NoCollision(t *testing.T) { + current := []string{"009a1f293253e41e"} + newIds := []string{"0039ff30789bc776"} + err := CheckCollidingKeysets(current, newIds) + if err != nil { + t.Fatalf("expected no collision, got: %v", err) + } +} + +func TestCheckCollidingKeysets_ExactMatch(t *testing.T) { + current := []string{"009a1f293253e41e"} + newIds := []string{"009a1f293253e41e"} + err := CheckCollidingKeysets(current, newIds) + if err == nil { + t.Fatal("expected collision error for exact match") + } +} + +func TestCheckCollidingKeysets_ModuloCollision(t *testing.T) { + id1 := "009a1f293253e41e" + bigId := "01df97b6fb8a572a718d7df7fcbf4387e2d455134ea8004c9c8c51e1b3391f909e" + + r1, _ := keysetIdToBigInt(id1) + r2, _ := keysetIdToBigInt(bigId) + if r1.Cmp(r2) == 0 { + err := CheckCollidingKeysets([]string{id1}, []string{bigId}) + if err == nil { + t.Fatal("expected collision for modulo-equal keysets") + } + } +}