Finish up Noise tests

This commit is contained in:
Nadim Kobeissi
2025-11-18 21:31:55 +02:00
parent cf528b0daf
commit f37acf7e4b
3 changed files with 77 additions and 33 deletions
+5 -5
View File
@@ -451,13 +451,13 @@ final class NoiseSymmetricState {
} }
} }
func split() -> (NoiseCipherState, NoiseCipherState) { func split(useExtractedNonce: Bool) -> (NoiseCipherState, NoiseCipherState) {
let output = hkdf(chainingKey: chainingKey, inputKeyMaterial: Data(), numOutputs: 2) let output = hkdf(chainingKey: chainingKey, inputKeyMaterial: Data(), numOutputs: 2)
let tempKey1 = SymmetricKey(data: output[0]) let tempKey1 = SymmetricKey(data: output[0])
let tempKey2 = SymmetricKey(data: output[1]) let tempKey2 = SymmetricKey(data: output[1])
let c1 = NoiseCipherState(key: tempKey1, useExtractedNonce: true) let c1 = NoiseCipherState(key: tempKey1, useExtractedNonce: useExtractedNonce)
let c2 = NoiseCipherState(key: tempKey2, useExtractedNonce: true) let c2 = NoiseCipherState(key: tempKey2, useExtractedNonce: useExtractedNonce)
return (c1, c2) return (c1, c2)
} }
@@ -789,12 +789,12 @@ final class NoiseHandshakeState {
return currentPattern >= messagePatterns.count return currentPattern >= messagePatterns.count
} }
func getTransportCiphers() throws -> (send: NoiseCipherState, receive: NoiseCipherState) { func getTransportCiphers(useExtractedNonce: Bool) throws -> (send: NoiseCipherState, receive: NoiseCipherState) {
guard isHandshakeComplete() else { guard isHandshakeComplete() else {
throw NoiseError.handshakeNotComplete throw NoiseError.handshakeNotComplete
} }
let (c1, c2) = symmetricState.split() let (c1, c2) = symmetricState.split(useExtractedNonce: useExtractedNonce)
// Initiator uses c1 for sending, c2 for receiving // Initiator uses c1 for sending, c2 for receiving
// Responder uses c2 for sending, c1 for receiving // Responder uses c2 for sending, c1 for receiving
+2 -2
View File
@@ -103,7 +103,7 @@ class NoiseSession {
// Check if handshake is complete // Check if handshake is complete
if handshake.isHandshakeComplete() { if handshake.isHandshakeComplete() {
// Get transport ciphers // Get transport ciphers
let (send, receive) = try handshake.getTransportCiphers() let (send, receive) = try handshake.getTransportCiphers(useExtractedNonce: true)
sendCipher = send sendCipher = send
receiveCipher = receive receiveCipher = receive
@@ -129,7 +129,7 @@ class NoiseSession {
// Check if handshake is complete after writing // Check if handshake is complete after writing
if handshake.isHandshakeComplete() { if handshake.isHandshakeComplete() {
// Get transport ciphers // Get transport ciphers
let (send, receive) = try handshake.getTransportCiphers() let (send, receive) = try handshake.getTransportCiphers(useExtractedNonce: true)
sendCipher = send sendCipher = send
receiveCipher = receive receiveCipher = receive
+70 -26
View File
@@ -679,23 +679,59 @@ struct NoiseProtocolTests {
predeterminedEphemeralKey: respEphemeralKey predeterminedEphemeralKey: respEphemeralKey
) )
// Message 1: Initiator -> Responder (e) // For XX pattern, we have 3 handshake messages, then transport messages
let msg1 = try initiatorHandshake.writeMessage() // The test vector messages are ordered as: [msg1, msg2, msg3, transport1, transport2, ...]
#expect(!msg1.isEmpty, "Message 1 should not be empty")
_ = try responderHandshake.readMessage(msg1) guard testVector.messages.count >= 3 else {
throw NSError(
domain: "NoiseTests", code: 5,
userInfo: [NSLocalizedDescriptionKey: "Test vector must have at least 3 messages for XX pattern"])
}
// Message 1: Initiator -> Responder (e)
guard let payload1 = Data(hex: testVector.messages[0].payload),
let expectedCiphertext1 = Data(hex: testVector.messages[0].ciphertext) else {
throw NSError(
domain: "NoiseTests", code: 4,
userInfo: [NSLocalizedDescriptionKey: "Message 1: Failed to parse hex"])
}
let msg1 = try initiatorHandshake.writeMessage(payload: payload1)
#expect(!msg1.isEmpty, "Message 1 should not be empty")
#expect(msg1 == expectedCiphertext1, "Message 1 ciphertext should match expected value. Got: \(msg1.hexString()), Expected: \(expectedCiphertext1.hexString())")
let decrypted1 = try responderHandshake.readMessage(msg1)
#expect(decrypted1 == payload1, "Message 1: Decrypted payload should match original")
// Message 2: Responder -> Initiator (e, ee, s, es) // Message 2: Responder -> Initiator (e, ee, s, es)
let msg2 = try responderHandshake.writeMessage() guard let payload2 = Data(hex: testVector.messages[1].payload),
#expect(!msg2.isEmpty, "Message 2 should not be empty") let expectedCiphertext2 = Data(hex: testVector.messages[1].ciphertext) else {
throw NSError(
domain: "NoiseTests", code: 4,
userInfo: [NSLocalizedDescriptionKey: "Message 2: Failed to parse hex"])
}
_ = try initiatorHandshake.readMessage(msg2) let msg2 = try responderHandshake.writeMessage(payload: payload2)
#expect(!msg2.isEmpty, "Message 2 should not be empty")
#expect(msg2 == expectedCiphertext2, "Message 2 ciphertext should match expected value. Got: \(msg2.hexString()), Expected: \(expectedCiphertext2.hexString())")
let decrypted2 = try initiatorHandshake.readMessage(msg2)
#expect(decrypted2 == payload2, "Message 2: Decrypted payload should match original")
// Message 3: Initiator -> Responder (s, se) // Message 3: Initiator -> Responder (s, se)
let msg3 = try initiatorHandshake.writeMessage() guard let payload3 = Data(hex: testVector.messages[2].payload),
#expect(!msg3.isEmpty, "Message 3 should not be empty") let expectedCiphertext3 = Data(hex: testVector.messages[2].ciphertext) else {
throw NSError(
domain: "NoiseTests", code: 4,
userInfo: [NSLocalizedDescriptionKey: "Message 3: Failed to parse hex"])
}
_ = try responderHandshake.readMessage(msg3) let msg3 = try initiatorHandshake.writeMessage(payload: payload3)
#expect(!msg3.isEmpty, "Message 3 should not be empty")
#expect(msg3 == expectedCiphertext3, "Message 3 ciphertext should match expected value. Got: \(msg3.hexString()), Expected: \(expectedCiphertext3.hexString())")
let decrypted3 = try responderHandshake.readMessage(msg3)
#expect(decrypted3 == payload3, "Message 3: Decrypted payload should match original")
// Verify handshake hash // Verify handshake hash
let initiatorHash = initiatorHandshake.getHandshakeHash() let initiatorHash = initiatorHandshake.getHandshakeHash()
@@ -706,16 +742,18 @@ struct NoiseProtocolTests {
if let expectedHash = expectedHash { if let expectedHash = expectedHash {
#expect( #expect(
initiatorHash == expectedHash, initiatorHash == expectedHash,
"Handshake hash should match expected value from test vector") "Handshake hash should match expected value from test vector. Got: \(initiatorHash.hexString()), Expected: \(expectedHash.hexString())")
} }
// Get transport ciphers // Get transport ciphers
let (initSend, initRecv) = try initiatorHandshake.getTransportCiphers() let (initSend, initRecv) = try initiatorHandshake.getTransportCiphers(useExtractedNonce: false)
let (respSend, respRecv) = try responderHandshake.getTransportCiphers() let (respSend, respRecv) = try responderHandshake.getTransportCiphers(useExtractedNonce: false)
// Test transport messages // Test transport messages (messages after the 3 handshake messages)
for (index, testMsg) in testVector.messages.enumerated() { for index in 3..<testVector.messages.count {
guard let payload = Data(hex: testMsg.payload) else { let testMsg = testVector.messages[index]
guard let payload = Data(hex: testMsg.payload),
let expectedCiphertext = Data(hex: testMsg.ciphertext) else {
throw NSError( throw NSError(
domain: "NoiseTests", code: 4, domain: "NoiseTests", code: 4,
userInfo: [ userInfo: [
@@ -724,22 +762,28 @@ struct NoiseProtocolTests {
]) ])
} }
// Alternate between initiator and responder sending // Alternate between responder and initiator sending
// Responder sends first transport message (since initiator sent last handshake message)
let (sender, receiver): (NoiseCipherState, NoiseCipherState) let (sender, receiver): (NoiseCipherState, NoiseCipherState)
if index % 2 == 0 { let transportIndex = index - 3
sender = initSend if transportIndex % 2 == 0 {
receiver = respRecv // Even transport messages: responder sends
} else {
sender = respSend sender = respSend
receiver = initRecv receiver = initRecv
} else {
// Odd transport messages: initiator sends
sender = initSend
receiver = respRecv
} }
// Encrypt // Encrypt and validate ciphertext matches expected value
let ciphertext = try sender.encrypt(plaintext: payload) let ciphertext = try sender.encrypt(plaintext: payload)
#expect(
// Decrypt ciphertext == expectedCiphertext,
"Message \(index + 1) ciphertext should match expected value. Got: \(ciphertext.hexString()), Expected: \(expectedCiphertext.hexString())")
// Decrypt and validate payload
let decrypted = try receiver.decrypt(ciphertext: ciphertext) let decrypted = try receiver.decrypt(ciphertext: ciphertext)
#expect( #expect(
decrypted == payload, decrypted == payload,
"Message \(index + 1): Decrypted payload should match original") "Message \(index + 1): Decrypted payload should match original")