mirror of
https://github.com/permissionlesstech/bitchat.git
synced 2026-07-25 15:25:19 +00:00
Finish up Noise tests
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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(
|
||||||
|
ciphertext == expectedCiphertext,
|
||||||
|
"Message \(index + 1) ciphertext should match expected value. Got: \(ciphertext.hexString()), Expected: \(expectedCiphertext.hexString())")
|
||||||
|
|
||||||
// Decrypt
|
// 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")
|
||||||
|
|||||||
Reference in New Issue
Block a user