diff --git a/include/sframe/result.h b/include/sframe/result.h index 0a74c20..18e5edc 100644 --- a/include/sframe/result.h +++ b/include/sframe/result.h @@ -1,5 +1,6 @@ #pragma once +#include #include #include #include @@ -18,6 +19,7 @@ enum class SFrameErrorType unsupported_ciphersuite_error, authentication_error, invalid_key_usage_error, + unknown_key_id_error, }; class SFrameError @@ -35,6 +37,13 @@ class SFrameError { } + SFrameError(SFrameErrorType type, const char* message, uint64_t key_id) + : type_(type) + , message_(message) + , key_id_(key_id) + { + } + SFrameError(const SFrameError& other) = default; SFrameError(SFrameError&& other) noexcept = default; SFrameError& operator=(SFrameError&& other) noexcept = default; @@ -43,11 +52,15 @@ class SFrameError const char* message() const { return message_; } + // Populated only when type() == SFrameErrorType::unknown_key_id_error. + std::optional key_id() const { return key_id_; } + private: SFrameErrorType type_; // Message storage is borrowed; callers must pass a string with static or // otherwise stable lifetime. const char* message_ = nullptr; + std::optional key_id_ = std::nullopt; }; #ifdef __cpp_exceptions diff --git a/include/sframe/sframe.h b/include/sframe/sframe.h index 397d03e..13ccb2d 100644 --- a/include/sframe/sframe.h +++ b/include/sframe/sframe.h @@ -65,6 +65,12 @@ struct invalid_key_usage_error : std::runtime_error using parent = std::runtime_error; using parent::parent; }; + +struct unknown_key_id_error : std::runtime_error +{ + using parent = std::runtime_error; + using parent::parent; +}; #endif enum class CipherSuite : uint16_t diff --git a/src/result.cpp b/src/result.cpp index 6fae585..3092125 100644 --- a/src/result.cpp +++ b/src/result.cpp @@ -32,6 +32,8 @@ throw_sframe_error(const SFrameError& error) throw authentication_error(); case SFrameErrorType::invalid_key_usage_error: throw invalid_key_usage_error(error.message()); + case SFrameErrorType::unknown_key_id_error: + throw unknown_key_id_error(error.message()); } } #endif diff --git a/src/sframe.cpp b/src/sframe.cpp index 60c28c1..3fbf77f 100644 --- a/src/sframe.cpp +++ b/src/sframe.cpp @@ -130,8 +130,8 @@ Result Context::require_key(KeyID key_id) const { if (!keys.contains(key_id)) { - return SFrameError(SFrameErrorType::invalid_parameter_error, - "Unknown key ID"); + return SFrameError( + SFrameErrorType::unknown_key_id_error, "Unknown key ID", key_id); } return Result::ok(); } @@ -419,8 +419,8 @@ MLSContext::ensure_key(KeyID key_id, KeyUsage usage) const auto epoch_index = key_id & epoch_mask; auto& epoch = epoch_cache[epoch_index]; if (!epoch) { - return SFrameError(SFrameErrorType::invalid_parameter_error, - "Unknown epoch"); + return SFrameError( + SFrameErrorType::unknown_key_id_error, "Unknown key ID", key_id); } if (keys.contains(key_id)) { diff --git a/test/sframe.cpp b/test/sframe.cpp index 9f8adf8..146bf59 100644 --- a/test/sframe.cpp +++ b/test/sframe.cpp @@ -223,7 +223,7 @@ TEST_CASE("MLS Failure after Purge") .error() .type() == SFrameErrorType::invalid_parameter_error); CHECK(member_b.unprotect(pt_out, enc_ab_1_data, metadata).error().type() == - SFrameErrorType::invalid_parameter_error); + SFrameErrorType::unknown_key_id_error); const auto enc_ab_2 = member_a.protect(epoch_id_2, sender_id_a, ct_out, plaintext, metadata) @@ -258,12 +258,12 @@ TEST_CASE("SFrame Context Remove Key") // Remove sender key and verify protect fails sender.remove_key(kid); CHECK(sender.protect(kid, ct_out, plaintext, metadata).error().type() == - SFrameErrorType::invalid_parameter_error); + SFrameErrorType::unknown_key_id_error); // Remove receiver key and verify unprotect fails receiver.remove_key(kid); CHECK(receiver.unprotect(pt_out, encrypted, metadata).error().type() == - SFrameErrorType::invalid_parameter_error); + SFrameErrorType::unknown_key_id_error); // Re-add keys and verify round-trip works again sender.add_key(kid, KeyUsage::protect, key).unwrap(); @@ -286,6 +286,73 @@ TEST_CASE("SFrame Context Remove Key - Nonexistent Key") CHECK_NOTHROW(ctx.remove_key(KeyID(0x99))); } +TEST_CASE("SFrame Unknown Key") +{ + const auto suite = CipherSuite::AES_GCM_128_SHA256; + const auto kid = KeyID(0x42); + const auto unknown_kid = KeyID(0x43); + const auto key = from_hex("000102030405060708090a0b0c0d0e0f"); + const auto plaintext = from_hex("00010203"); + const auto metadata = bytes{}; + + auto pt_out = bytes(plaintext.size()); + auto ct_out = bytes(plaintext.size() + Context::max_overhead); + + auto sender = Context(suite); + sender.add_key(kid, KeyUsage::protect, key).unwrap(); + + // Protecting with a key ID that was never added fails with + // unknown_key_id_error + CHECK( + sender.protect(unknown_kid, ct_out, plaintext, metadata).error().type() == + SFrameErrorType::unknown_key_id_error); + + // Produce a valid ciphertext whose header references `kid` + auto encrypted = + to_bytes(sender.protect(kid, ct_out, plaintext, metadata).unwrap()); + + // A receiver that doesn't know `kid` fails to unprotect with + // unknown_key_id_error + auto receiver = Context(suite); + CHECK(receiver.unprotect(pt_out, encrypted, metadata).error().type() == + SFrameErrorType::unknown_key_id_error); +} + +TEST_CASE("SFrame Unknown Key Error Reports Key ID") +{ + const auto suite = CipherSuite::AES_GCM_128_SHA256; + const auto kid = KeyID(0x42); + const auto unknown_kid = KeyID(0x43); + const auto key = from_hex("000102030405060708090a0b0c0d0e0f"); + const auto plaintext = from_hex("00010203"); + const auto metadata = bytes{}; + + auto pt_out = bytes(plaintext.size()); + auto ct_out = bytes(plaintext.size() + Context::max_overhead); + + auto sender = Context(suite); + sender.add_key(kid, KeyUsage::protect, key).unwrap(); + + // Protecting with a key ID never added: error must name that key ID + auto protect_err = + sender.protect(unknown_kid, ct_out, plaintext, metadata).error(); + CHECK(protect_err.type() == SFrameErrorType::unknown_key_id_error); + CHECK(protect_err.key_id().has_value()); + CHECK(protect_err.key_id().value() == unknown_kid); + + // Produce a ciphertext whose header embeds kid + auto encrypted = + to_bytes(sender.protect(kid, ct_out, plaintext, metadata).unwrap()); + + // Receiver with no keys: error must name the key ID parsed from the + // ciphertext + auto receiver = Context(suite); + auto unprotect_err = receiver.unprotect(pt_out, encrypted, metadata).error(); + CHECK(unprotect_err.type() == SFrameErrorType::unknown_key_id_error); + CHECK(unprotect_err.key_id().has_value()); + CHECK(unprotect_err.key_id().value() == kid); +} + TEST_CASE("MLS Remove Epoch") { const auto suite = CipherSuite::AES_GCM_128_SHA256; @@ -327,7 +394,7 @@ TEST_CASE("MLS Remove Epoch") .error() .type() == SFrameErrorType::invalid_parameter_error); CHECK(member_b.unprotect(pt_out, enc_data, metadata).error().type() == - SFrameErrorType::invalid_parameter_error); + SFrameErrorType::unknown_key_id_error); // Epoch 2 should still work enc = member_a.protect(epoch_id_2, sender_id, ct_out, plaintext, metadata)