diff --git a/bitchat/Noise/NoiseProtocol.swift b/bitchat/Noise/NoiseProtocol.swift index df92bae8..021dc9fc 100644 --- a/bitchat/Noise/NoiseProtocol.swift +++ b/bitchat/Noise/NoiseProtocol.swift @@ -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 tempKey1 = SymmetricKey(data: output[0]) let tempKey2 = SymmetricKey(data: output[1]) - let c1 = NoiseCipherState(key: tempKey1, useExtractedNonce: true) - let c2 = NoiseCipherState(key: tempKey2, useExtractedNonce: true) + let c1 = NoiseCipherState(key: tempKey1, useExtractedNonce: useExtractedNonce) + let c2 = NoiseCipherState(key: tempKey2, useExtractedNonce: useExtractedNonce) return (c1, c2) } @@ -789,12 +789,12 @@ final class NoiseHandshakeState { return currentPattern >= messagePatterns.count } - func getTransportCiphers() throws -> (send: NoiseCipherState, receive: NoiseCipherState) { + func getTransportCiphers(useExtractedNonce: Bool) throws -> (send: NoiseCipherState, receive: NoiseCipherState) { guard isHandshakeComplete() else { throw NoiseError.handshakeNotComplete } - let (c1, c2) = symmetricState.split() + let (c1, c2) = symmetricState.split(useExtractedNonce: useExtractedNonce) // Initiator uses c1 for sending, c2 for receiving // Responder uses c2 for sending, c1 for receiving diff --git a/bitchat/Noise/NoiseSession.swift b/bitchat/Noise/NoiseSession.swift index 52e3dedd..8c84f85a 100644 --- a/bitchat/Noise/NoiseSession.swift +++ b/bitchat/Noise/NoiseSession.swift @@ -103,7 +103,7 @@ class NoiseSession { // Check if handshake is complete if handshake.isHandshakeComplete() { // Get transport ciphers - let (send, receive) = try handshake.getTransportCiphers() + let (send, receive) = try handshake.getTransportCiphers(useExtractedNonce: true) sendCipher = send receiveCipher = receive @@ -129,7 +129,7 @@ class NoiseSession { // Check if handshake is complete after writing if handshake.isHandshakeComplete() { // Get transport ciphers - let (send, receive) = try handshake.getTransportCiphers() + let (send, receive) = try handshake.getTransportCiphers(useExtractedNonce: true) sendCipher = send receiveCipher = receive diff --git a/bitchatTests/Noise/NoiseProtocolTests.swift b/bitchatTests/Noise/NoiseProtocolTests.swift index 3b790642..6083e134 100644 --- a/bitchatTests/Noise/NoiseProtocolTests.swift +++ b/bitchatTests/Noise/NoiseProtocolTests.swift @@ -679,23 +679,59 @@ struct NoiseProtocolTests { predeterminedEphemeralKey: respEphemeralKey ) - // Message 1: Initiator -> Responder (e) - let msg1 = try initiatorHandshake.writeMessage() - #expect(!msg1.isEmpty, "Message 1 should not be empty") + // For XX pattern, we have 3 handshake messages, then transport messages + // The test vector messages are ordered as: [msg1, msg2, msg3, transport1, transport2, ...] - _ = 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) - let msg2 = try responderHandshake.writeMessage() - #expect(!msg2.isEmpty, "Message 2 should not be empty") + guard let payload2 = Data(hex: testVector.messages[1].payload), + 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) - let msg3 = try initiatorHandshake.writeMessage() - #expect(!msg3.isEmpty, "Message 3 should not be empty") + guard let payload3 = Data(hex: testVector.messages[2].payload), + 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 let initiatorHash = initiatorHandshake.getHandshakeHash() @@ -706,16 +742,18 @@ struct NoiseProtocolTests { if let expectedHash = expectedHash { #expect( 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 - let (initSend, initRecv) = try initiatorHandshake.getTransportCiphers() - let (respSend, respRecv) = try responderHandshake.getTransportCiphers() - - // Test transport messages - for (index, testMsg) in testVector.messages.enumerated() { - guard let payload = Data(hex: testMsg.payload) else { + let (initSend, initRecv) = try initiatorHandshake.getTransportCiphers(useExtractedNonce: false) + let (respSend, respRecv) = try responderHandshake.getTransportCiphers(useExtractedNonce: false) + + // Test transport messages (messages after the 3 handshake messages) + for index in 3..