diff --git a/ecdsa.go b/ecdsa.go index 841486c1..b7d2c848 100644 --- a/ecdsa.go +++ b/ecdsa.go @@ -58,19 +58,26 @@ func bitsToBytes(bits int) int { return (bits + 7) >> 3 } -func (a *ecdsaContext) checkAlgoAndComputeHash(msg []byte, hasher hash.Hasher) (hash.Hash, error) { +// checkHasherAndComputeHash checks the hasher is valid for ECDSA +// on the receiver curve and returns the hash of the input message. +func (a *ecdsaContext) checkHasherAndComputeHash(msg []byte, hasher hash.Hasher) (hash.Hash, error) { if hasher == nil { return nil, errNilHasher } - // check hasher's size is at least the curve order in bytes - nLen := (a.curveN).BitLen() - if (hasher.Size() << 3) < nLen { + h := hasher.ComputeHash(msg) + // check the computed hash is at least the curve order in bytes. + // All curve orders supported by the package have a bit-length multiple of 8, + // so callers truncate the message hash in bytes + // and the check is done in bytes too. + // The check uses the computed hash length rather than the hasher's declared size, + // so that a hasher implementation computing fewer bytes than it declares + // is rejected instead of panicking in the caller's truncation. + nLen := bitsToBytes((a.curveN).BitLen()) + if len(h) < nLen { return nil, invalidHasherSizeErrorf( - "hasher's bit-size should be at least %d, got %d", nLen, hasher.Size()<<3) + "hasher's output should be at least %d bytes, got %d bytes", nLen, len(h)) } - - h := hasher.ComputeHash(msg) return h, nil } @@ -106,7 +113,6 @@ func (a *ecdsaContext) signatureFormatCheck(sig Signature) bool { } var one = new(big.Int).SetInt64(1) -var two = new(big.Int).SetInt64(2) // mapToPrivateKey simply maps the input seed to an ECDSA private key // The private scalar `d` satisfies 0 < d < n. @@ -137,14 +143,23 @@ func (a *ecdsaContext) privateKey(d *big.Int) (PrivateKey, error) { d.FillBytes(dBytes) // dBytes is the big-endian encoding of d padded to the curve order // build the private key depending on the curve + var sk PrivateKey + var err error switch a.algo { case ECDSAP256: - return privateKeyECDSAP256(a, dBytes) + sk, err = privateKeyECDSAP256(a, dBytes) case ECDSASecp256k1: - return privateKeyECDSASecp256k1(a, dBytes), nil + sk = privateKeyECDSASecp256k1(a, dBytes) default: return nil, invalidInputsErrorf("the curve is not supported") } + if err != nil { + // return an untyped nil, + // otherwise the returned interface is non-nil + // although it holds a nil pointer + return nil, err + } + return sk, nil } // generatePrivateKey generates a private key for ECDSA @@ -219,14 +234,23 @@ func (a *ecdsaContext) decodePrivateKey(der []byte) (PrivateKey, error) { // Error Returns: // - invalidInputsError if the input is not a valid serialization of a public key on the given curve. func (a *ecdsaContext) rawDecodePublicKey(input []byte) (PublicKey, error) { + var pk PublicKey + var err error switch a.algo { case ECDSAP256: - return publicKeyECDSAP256(input) + pk, err = publicKeyECDSAP256(input) case ECDSASecp256k1: - return publicKeyECDSASecp256k1(a, input) + pk, err = publicKeyECDSASecp256k1(a, input) default: return nil, invalidInputsErrorf("curve is not supported") } + if err != nil { + // return an untyped nil, + // otherwise the returned interface is non-nil + // although it holds a nil pointer + return nil, err + } + return pk, nil } func (a *ecdsaContext) decodePublicKey(der []byte) (PublicKey, error) { @@ -241,14 +265,23 @@ func (a *ecdsaContext) decodePublicKey(der []byte) (PublicKey, error) { // - invalidInputsError if the curve isn't supported or the input isn't a valid key serialization // on the given curve. func (a *ecdsaContext) decodePublicKeyCompressed(pkBytes []byte) (PublicKey, error) { + var pk PublicKey + var err error switch a.algo { case ECDSAP256: - return p256DecodePublicKeyCompressed(pkBytes) + pk, err = p256DecodePublicKeyCompressed(pkBytes) case ECDSASecp256k1: - return secp256k1DecodePublicKeyCompressed(pkBytes) + pk, err = secp256k1DecodePublicKeyCompressed(pkBytes) default: return nil, invalidInputsErrorf("the input curve is not supported") } + if err != nil { + // return an untyped nil, + // otherwise the returned interface is non-nil + // although it holds a nil pointer + return nil, err + } + return pk, nil } // Algorithm returns the algo related to the private key @@ -280,6 +313,10 @@ func pubKeyCommonECDSAString(pk PublicKey) string { // Equals tests the equality of two private keys func prKeyCommonECDSAEquals(sk, other PrivateKey) bool { + // a nil key is not equal to any key + if other == nil { + return false + } // check the algorithm if sk.Algorithm() != other.Algorithm() { return false @@ -300,6 +337,10 @@ func (pk *pubKeyCommonECDSA) Size() int { // Equals tests the equality of two private keys func pubKeyCommonECDSAEquals(pk, other PublicKey) bool { + // a nil key is not equal to any key + if other == nil { + return false + } // check the algorithm if pk.Algorithm() != other.Algorithm() { return false @@ -313,7 +354,7 @@ func pubKeyCommonECDSAEquals(pk, other PublicKey) bool { // It assumes the output buffer has at least 2*size byte-length func padToSizeAndConcat(output []byte, a, b *big.Int, size int) { a.FillBytes(output[:size]) - b.FillBytes(output[size:]) + b.FillBytes(output[size : 2*size]) } // Helper function to read two big integers of "size" bytes each from a concatenated input buffer. @@ -333,7 +374,9 @@ func (a *ecdsaContext) isLowS(s *big.Int) bool { // signatureNormalizeS returns a signature with S normalized to low S. // (same slice is returned if S is already normalized) // It assumes len(sig) == 2*nLen where nLen is the byte-length of the curve order. -// This is needed when the underlying signature verification requires S to be in the lower range (to avoid signature malleability). In this package, verification allows high S signatures to be accepted. +// This is needed when the underlying signature verification requires S to be +// in the lower range (to avoid signature malleability). +// In this package, verification allows high S signatures to be accepted. // The function checks that S is in the range [0, n-1] before normalizing it. // If S is not in this range, the function returns a false boolean. // (S and R values will be checked by the go-ethereum verification function - only S check against N is included here, S=0 check is deferred to the signature verification) @@ -361,17 +404,3 @@ func (a *ecdsaContext) signatureNormalizeS(sig []byte) ([]byte, bool) { sComplement.FillBytes(newSig[nLen:]) // write S complement return newSig, true } - -// Test function only to flip S in a signature. It is used for testing signature malleability -func (a *ecdsaContext) signatureFlipS(sig []byte) []byte { - // read S - nLen := bitsToBytes(a.curveN.BitLen()) - s := new(big.Int).SetBytes(sig[nLen:]) - // compute N-S - sComplement := new(big.Int).Sub(a.curveN, s) - // write it into a new signature - newSig := make([]byte, len(sig)) - copy(newSig, sig[:nLen]) // copy R - sComplement.FillBytes(newSig[nLen:]) // write S complement - return newSig -} diff --git a/ecdsa_p256.go b/ecdsa_p256.go index a5b8e02f..8f6ade9e 100644 --- a/ecdsa_p256.go +++ b/ecdsa_p256.go @@ -50,12 +50,11 @@ var p256Instance *ecdsaContext func initECDSAP256() { curve := elliptic.P256() n := curve.Params().N - nMinus1 := new(big.Int).Sub(n, one) p256Instance = &(ecdsaContext{ curveP: curve.Params().P, curveN: n, - curveNdiv2: new(big.Int).Div(nMinus1, two), // (N-1)/2 + curveNdiv2: new(big.Int).Rsh(n, 1), // (N-1)/2, since N is odd algo: ECDSAP256, }) } @@ -91,7 +90,10 @@ func privateKeyECDSAP256(a *ecdsaContext, dBytes []byte) (*prKeyECDSAP256, error sk := &prKeyECDSAP256{ prKeyCommonECDSA: &prKeyCommonECDSA{a}, goPrKey: internalSK, - pubKey: nil, // public key is not constructed yet + pubKey: &pubKeyECDSAP256{ + pubKeyCommonECDSA: &pubKeyCommonECDSA{p256Instance}, + goPubKey: &internalSK.PublicKey, + }, } return sk, nil } @@ -109,7 +111,7 @@ func privateKeyECDSAP256(a *ecdsaContext, dBytes []byte) (*prKeyECDSAP256, error // - (nil, error) if an unexpected error occurs // - (signature, nil) otherwise func (sk *prKeyECDSAP256) Sign(msg []byte, hasher hash.Hasher) (Signature, error) { - hash, err := sk.checkAlgoAndComputeHash(msg, hasher) + hash, err := sk.checkHasherAndComputeHash(msg, hasher) if err != nil { return nil, err } @@ -135,10 +137,12 @@ func publicKeyECDSAP256(XYBytes []byte) (*pubKeyECDSAP256, error) { // The bytes serialization for non-infinity points is `0x04 || X || Y` and infinity point should be rejected anyway parsingBytes := append([]byte{ecEncodingUncompressed}, XYBytes...) - // ParseUncompressedPublicKey includes x