diff --git a/std/security/certificate.go b/std/security/certificate.go index a9dd4386..10698cf6 100644 --- a/std/security/certificate.go +++ b/std/security/certificate.go @@ -213,17 +213,29 @@ func revocationRecordName(cert ndn.Data) (enc.Name, bool) { return recordName.Append(certName.At(-2)), true } -func isRevocationRecordName(name enc.Name) bool { +// CertNameFromRevocationRecordName returns the certificate name encoded by a +// revocation record name. +func CertNameFromRevocationRecordName(name enc.Name) (enc.Name, error) { name = stripImplicitDigest(name) - if len(name) < 6 || name.At(-1).Typ != enc.TypeGenericNameComponent || !name.At(-2).IsVersion() { - return false + if len(name) < 6 || + !name.At(-5).IsGeneric("REVOKE") || + !name.At(-2).IsVersion() || + !name.At(-1).Equal(name.At(-3)) { + return nil, fmt.Errorf("invalid revocation record name: %s", name) } certName := make(enc.Name, len(name)-1) copy(certName, name.Prefix(-1)) certName[len(name)-5] = enc.NewGenericComponent("KEY") - _, err := GetIdentityFromCertName(certName) + if _, err := GetIdentityFromCertName(certName); err != nil { + return nil, fmt.Errorf("invalid revocation record name %s: %w", name, err) + } + return certName, nil +} + +func isRevocationRecordName(name enc.Name) bool { + _, err := CertNameFromRevocationRecordName(name) return err == nil } diff --git a/std/security/certificate_test.go b/std/security/certificate_test.go index 9002e244..aa0f3e92 100644 --- a/std/security/certificate_test.go +++ b/std/security/certificate_test.go @@ -205,6 +205,26 @@ func TestCertificateRevocationRecord(t *testing.T) { require.True(t, ok) require.Equal(t, ndn.ContentTypeKey, contentType) require.NotEmpty(t, recordData.Content()) + + certName, err := sec.CertNameFromRevocationRecordName(recordData.Name()) + require.NoError(t, err) + require.True(t, aliceCert.Name().Equal(certName)) +} + +func TestCertNameFromRevocationRecordNameRejectsMalformedNames(t *testing.T) { + tests := []string{ + "/REVOKE/key/issuer/v=1/issuer", + "/test/alice/REVOKE/key/issuer/not-version/issuer", + "/test/alice/REVOKE/key/issuer/v=1/v=2", + "/test/alice/KEY/key/issuer/v=1/issuer", + } + + for _, value := range tests { + t.Run(value, func(t *testing.T) { + _, err := sec.CertNameFromRevocationRecordName(tu.NoErr(enc.NameFromStr(value))) + require.Error(t, err) + }) + } } func revocationTestIssuer(t *testing.T) enc.Component { diff --git a/std/security/trust_config.go b/std/security/trust_config.go index dfaa22f1..3ee87cc3 100644 --- a/std/security/trust_config.go +++ b/std/security/trust_config.go @@ -1,6 +1,8 @@ package security import ( + "bytes" + "crypto/sha256" "fmt" "sync" @@ -107,13 +109,55 @@ func (tc *TrustConfig) InsertRevoke(wire enc.Wire) error { if !isRevocationRecordName(data.Name()) { return fmt.Errorf("not a revocation record name: %s", data.Name()) } + record, err := revocationtlv.ParseRevocationRecord(enc.NewWireView(data.Content()), false) + if err != nil { + return fmt.Errorf("failed to parse revocation record content: %w", err) + } + if len(record.PublicKeyHash) != sha256.Size { + return fmt.Errorf("invalid revocation public-key hash length: %d", len(record.PublicKeyHash)) + } tc.mutex.Lock() defer tc.mutex.Unlock() + if current, _ := tc.keychain.Store().Get(data.Name(), false); len(current) > 0 { + if bytes.Equal(current, wire.Join()) { + return nil + } + return fmt.Errorf("conflicting revocation record already stored: %s", data.Name()) + } // HACK: revocation records live in keychain.Store() until we have a dedicated store. return tc.keychain.Store().Put(data.Name(), wire.Join()) } +// CheckRevoke returns the stored revocation record for cert. A nil record +// means the certificate has no locally stored revocation. +func (tc *TrustConfig) CheckRevoke(cert ndn.Data) (*revocationtlv.RevocationRecord, error) { + if cert == nil { + return nil, fmt.Errorf("certificate is nil") + } + recordName, ok := revocationRecordName(cert) + if !ok { + return nil, fmt.Errorf("invalid certificate name: %s", cert.Name()) + } + + tc.mutex.RLock() + wire, err := tc.keychain.Store().Get(recordName, false) + tc.mutex.RUnlock() + if err != nil || len(wire) == 0 { + return nil, nil + } + + data, _, err := spec.Spec{}.ReadData(enc.NewBufferView(wire)) + if err != nil { + return nil, fmt.Errorf("failed to parse stored revocation record: %w", err) + } + record, err := revocationtlv.ParseRevocationRecord(enc.NewWireView(data.Content()), false) + if err != nil { + return nil, fmt.Errorf("failed to parse stored revocation content: %w", err) + } + return record, nil +} + // SetSchema atomically replaces the trust schema. func (tc *TrustConfig) SetSchema(schema ndn.TrustSchema) { if schema == nil { @@ -198,7 +242,8 @@ func (tc *TrustConfig) Validate(args TrustConfigValidateArgs) { // Bail if the data is a cert and is not fresh if t, ok := args.Data.ContentType().Get(); ok && t == ndn.ContentTypeKey { - if !args.IgnoreValidity.GetOr(false) && CertIsExpired(args.Data) { + if !isRevocationRecordName(args.Data.Name()) && + !args.IgnoreValidity.GetOr(false) && CertIsExpired(args.Data) { args.Callback(false, fmt.Errorf("certificate is expired: %s", args.Data.Name())) return } diff --git a/std/security/trust_config_test.go b/std/security/trust_config_test.go index 25044542..769dde6c 100644 --- a/std/security/trust_config_test.go +++ b/std/security/trust_config_test.go @@ -2,6 +2,7 @@ package security_test import ( "crypto/elliptic" + "crypto/sha256" _ "embed" "fmt" "math" @@ -1092,9 +1093,28 @@ func TestTrustConfigRevocation(t *testing.T) { require.Error(t, trust.InsertRevoke(nil)) require.Error(t, trust.InsertRevoke(aliceCertWire)) + _, err = trust.CheckRevoke(nil) + require.Error(t, err) recordWire := makeRevocationRecordWire(t, aliceCertData, rootSigner) require.NoError(t, trust.InsertRevoke(recordWire)) + require.NoError(t, trust.InsertRevoke(recordWire)) + conflictingRecord := tu.NoErr(sec.RevokeCert(sec.RevokeCertArgs{ + Cert: aliceCertData, + Signer: rootSigner, + Timestamp: optional.Some(time.Now().Add(time.Minute)), + })) + require.Error(t, trust.InsertRevoke(conflictingRecord)) + stored, err := trust.CheckRevoke(aliceCertData) + require.NoError(t, err) + require.NotNil(t, stored) + require.Equal(t, sha256.Sum256(aliceCertData.Content().Join()), [32]byte(stored.PublicKeyHash)) + + reloadedTrust, err := sec.NewTrustConfig(keychain, schema, []enc.Name{rootCertData.Name()}) + require.NoError(t, err) + stored, err = reloadedTrust.CheckRevoke(aliceCertData) + require.NoError(t, err) + require.NotNil(t, stored) recordData, _, err := spec.Spec{}.ReadData(enc.NewWireView(recordWire)) require.NoError(t, err)