diff --git a/crypto/key.go b/crypto/key.go index 1388ab7c..aac781ee 100644 --- a/crypto/key.go +++ b/crypto/key.go @@ -397,6 +397,31 @@ func (key *Key) GetSHA256Fingerprint() (fingerprint string) { return hex.EncodeToString(getSHA256FingerprintBytes(key.entity.PrimaryKey)) } +// IsForwardingKey checks if the given key is a Proton forwarding key. +func (key *Key) IsForwardingKey() bool { + decryptionKeys := key.entity.DecryptionKeys(0, time.Time{}, nil) + allForwarding := len(decryptionKeys) > 0 + for _, key := range decryptionKeys { + if !isForwardingKey(key) { + allForwarding = false + break + } + } + return allForwarding +} + +// isForwardingKey determines if the given openpgp.Key is a forwarding key. +func isForwardingKey(key openpgp.Key) bool { + curve, err := key.PublicKey.Curve() + keyCheck := key.PublicKey.IsSubkey && + key.PublicKey.Version == 4 && + key.PublicKey.PubKeyAlgo == packet.PubKeyAlgoECDH && + err == nil && curve == packet.Curve25519 + hasForwardFlag := key.SelfSignature != nil && + key.SelfSignature.FlagForward + return keyCheck && hasForwardFlag +} + // GetSHA256Fingerprints computes the SHA256 fingerprints of the key and subkeys. func (key *Key) GetSHA256Fingerprints() (fingerprints []string) { fingerprints = append(fingerprints, key.GetSHA256Fingerprint()) diff --git a/crypto/keyring.go b/crypto/keyring.go index 80a46fea..a8350e57 100644 --- a/crypto/keyring.go +++ b/crypto/keyring.go @@ -285,6 +285,19 @@ func FilterExpiredKeys(contactKeys []*KeyRing) (filteredKeys []*KeyRing, err err return filteredKeys, nil } +// WithoutForwardingKeys returns a new keyring containing only non-forwarding keys. +func (keyRing *KeyRing) WithoutForwardingKeys() (*KeyRing, error) { + filtered := &KeyRing{} + for _, key := range keyRing.GetKeys() { + if !key.IsForwardingKey() { + if err := filtered.AddKey(key); err != nil { + return nil, fmt.Errorf("gopenpgp: failed to add key to filtered keyring: %w", err) + } + } + } + return filtered, nil +} + // FirstKey returns a KeyRing with only the first key of the original one. func (keyRing *KeyRing) FirstKey() (*KeyRing, error) { if len(keyRing.entities) == 0 { diff --git a/crypto/proton_test.go b/crypto/proton_test.go index f608f5df..4e1ff8d0 100644 --- a/crypto/proton_test.go +++ b/crypto/proton_test.go @@ -36,6 +36,48 @@ func TestForwardeeDecryption(t *testing.T) { assert.Exactly(t, "Message for Bob", plainMessage.String()) } +func TestFowardingKeyCheck(t *testing.T) { + forwardingKey, err := NewKeyFromArmored(readTestFile("key_forwardee", false)) + if err != nil { + t.Fatal("Expected no error while unarmoring private keyring, got:", err) + } + + nonForwardingKey, err := NewKeyFromArmored(readTestFile("keyring_userKey", false)) + if err != nil { + t.Fatal("Expected no error while unarmoring private keyring, got:", err) + } + + if !forwardingKey.IsForwardingKey() { + t.Fatal("Expected a forwarding key") + } + + if nonForwardingKey.IsForwardingKey() { + t.Fatal("Expected non-forwarding key") + } + + kr, err := NewKeyRing(forwardingKey) + if err != nil { + t.Fatal(err) + } + + if err := kr.AddKey(nonForwardingKey); err != nil { + t.Fatal(err) + } + + krWithoutForwarding, err := kr.WithoutForwardingKeys() + if err != nil { + t.Fatal(err) + } + + assert.Exactly(t, krWithoutForwarding.CountEntities(), 1) + + key, err := krWithoutForwarding.GetKey(0) + if err != nil { + t.Fatal(err) + } + assert.True(t, !key.IsForwardingKey()) +} + func TestSymmetricKeys(t *testing.T) { symmetricKey, err := NewKeyFromArmored(readTestFile("key_symmetric", false)) if err != nil {