Repository navigation
TLS PSK intermittently fails with TlsFatalAlertReceived: bad_record_mac(20) #1876
Description
Activity
Test files: pom.xml and TlsPksTest.java
pom.xml
<?xml version="1.0" encoding="UTF-8"?> <!DOCTYPE RelativeLayout> <project xmlns="https://maven.apache.org/POM/4.0.0" xmlns:xsi="https://www.w3.org/2001/XMLSchema-instance" xsi:schemaLocation="https://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd"> <modelVersion>4.0.0</modelVersion> <groupId>com.github.justincranford</groupId> <artifactId>bc-tls-psk</artifactId> <version>0.0.1-SNAPSHOT</version> <name>BC TLS PSK</name> <description>BC TLS PSK</description> <dependencies> <dependency> <groupId>org.bouncycastle</groupId> <artifactId>bctls-debug-jdk18on</artifactId> <version>1.78.1</version> </dependency> <dependency> <groupId>org.slf4j</groupId> <artifactId>slf4j-simple</artifactId> <version>2.0.16</version> </dependency> <dependency> <groupId>org.assertj</groupId> <artifactId>assertj-core</artifactId> <version>3.26.3</version> <scope>test</scope> </dependency> <dependency> <groupId>org.junit.jupiter</groupId> <artifactId>junit-jupiter</artifactId> <version>5.11.3</version> <scope>test</scope> </dependency> <dependency> <groupId>org.mockito</groupId> <artifactId>mockito-core</artifactId> <version>5.14.2</version> <scope>test</scope> </dependency> </dependencies>
TlsPskTest.java
package com.github.justincranford.psk; import java.io.InputStream; import java.io.OutputStream; import java.net.InetAddress; import java.net.ServerSocket; import java.net.Socket; import java.nio.charset.StandardCharsets; import java.security.SecureRandom; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import org.assertj.core.api.Assertions; import org.bouncycastle.tls.CipherSuite; import org.bouncycastle.tls.PSKTlsClient; import org.bouncycastle.tls.PSKTlsServer; import org.bouncycastle.tls.TlsClientProtocol; import org.bouncycastle.tls.TlsPSKIdentity; import org.bouncycastle.tls.TlsPSKIdentityManager; import org.bouncycastle.tls.TlsServerProtocol; import org.bouncycastle.tls.crypto.impl.bc.BcTlsCrypto; import org.bouncycastle.util.io.Streams; import org.junit.jupiter.api.MethodOrderer; import org.junit.jupiter.api.Order; import org.junit.jupiter.api.TestMethodOrder; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.ValueSource; import org.mockito.Mockito; import org.slf4j.Logger; import org.slf4j.LoggerFactory; @TestMethodOrder(MethodOrderer.OrderAnnotation.class) @SuppressWarnings({"nls", "static-method", "hiding", "synthetic-access", "resource"}) public class PskTlsTest { private static final Logger log = LoggerFactory.getLogger(PskTlsTest.class); public static final SecureRandom SECURE_RANDOM = new SecureRandom(); private static final int[] CIPHER_SUITES = new int[] { CipherSuite.TLS_PSK_WITH_AES_128_CBC_SHA }; private static final TlsPskIdentity PSK_IDENTITY = new TlsPskIdentity("identity".getBytes(StandardCharsets.UTF_8),"secret".getBytes(StandardCharsets.UTF_8)); @ParameterizedTest // repeat test, use unique port each time to avoid TCP CLOSE_WAIT @ValueSource(ints={9440, 9441, 9442, 9443, 9444, 9445, 9446, 9447, 9448, 9449}) @Order(1) public void testPlaintext(final int port) throws Exception { doClientServer(false, "localhost", port); } @ParameterizedTest // repeat test, use unique port each time to avoid TCP CLOSE_WAIT @ValueSource(ints={8440, 8441, 8442, 8443, 8444, 8445, 8446, 8447, 8448, 8449}) @Order(2) public void testTlsPsk(final int port) throws Exception { doClientServer(true, "localhost", port); } // Start server, send message with client, and verify client received echo of its request // useTlsPsk=true uses plaintext communication // useTlsPsk=true uses plaintext communication private void doClientServer(final boolean useTlsPsk, final String address, final int port) throws Exception { final PskTlsServer pskTlsServer = Mockito.spy(new PskTlsServer(useTlsPsk, address, port, 2)); final String clientRequest = "This is an echo test " + SECURE_RANDOM.nextInt(); final Thread serverThread = pskTlsServer.start(); final String serverResponse = PskTlsClient.send(useTlsPsk, address, port, clientRequest); Assertions.assertThat(serverResponse).isEqualTo(clientRequest); serverThread.interrupt(); } public static class PskTlsClient { public static String send(final boolean useTlsPsk, final String address, final int port, final String clientRequest) throws Exception { log.info("Client: Connecting to server port " + port); try (final Socket socket = new Socket(address, port)) { log.info("Client: Connected to server port " + port); final InputStream inputStream = socket.getInputStream(); final OutputStream outputStream = socket.getOutputStream(); final byte[] serverResponseBytes = new byte[clientRequest.length()]; if (useTlsPsk) { // TLS PSK send and receive final TlsClientProtocol tlsClientProtocol = new TlsClientProtocol(inputStream, outputStream); tlsClientProtocol.connect(new PSKTlsClient(new BcTlsCrypto(SECURE_RANDOM), PSK_IDENTITY) { @Override public int[] getCipherSuites() { return CIPHER_SUITES; } }); final OutputStream tlsOutputStream = tlsClientProtocol.getOutputStream(); log.info("Client: Sending \"Hello from PSK Client\""); tlsOutputStream.write(clientRequest.getBytes(StandardCharsets.UTF_8)); tlsOutputStream.flush(); final InputStream tlsInputStream = tlsClientProtocol.getInputStream(); final int numServerResponseBytes = tlsInputStream.read(serverResponseBytes); Assertions.assertThat(numServerResponseBytes).isEqualTo(clientRequest.length()); tlsClientProtocol.close(); } else { // PLAINTEXT send and receive log.info("Client: Sending \"Hello from PSK Client\""); outputStream.write(clientRequest.getBytes(StandardCharsets.UTF_8)); outputStream.flush(); final int numServerResponseBytes = inputStream.read(serverResponseBytes); Assertions.assertThat(numServerResponseBytes).isEqualTo(clientRequest.length()); } final String serverResponse = new String(serverResponseBytes, StandardCharsets.UTF_8); log.info("Client: Received from server: " + serverResponse); return serverResponse; } } } public static class PskTlsServer { private static final int MAX_WAIT_MILLIS = 3000; private final boolean useTlsPsk; private final String address; private final int port; private final int backlog; public PskTlsServer(final boolean useTlsPsk, final String address, final int port, final int backlog) { this.useTlsPsk = useTlsPsk; this.address = address; this.port = port; this.backlog = backlog; } public void listen(final CountDownLatch countDownLatch) throws Exception { try (ServerSocket serverSocket = new ServerSocket(this.port, this.backlog, InetAddress.getByName(this.address))) { log.info("Server: Listening on " + this.address + ":" + this.port + "..."); countDownLatch.countDown(); // signal to main thread that server started OK while (true) { log.info("Server: While loop"); try (final Socket socket = serverSocket.accept()) { log.info("Server: Accepted connection from client"); final InputStream inputStream = socket.getInputStream(); final OutputStream outputStream = socket.getOutputStream(); if (this.useTlsPsk) { // TLS PSK echo final TlsServerProtocol tlsServerProtocol = new TlsServerProtocol(inputStream, outputStream); final BcTlsCrypto bcTlsCrypto = new BcTlsCrypto(SECURE_RANDOM); final PSKTlsServer pskTlsServer = new PSKTlsServer(bcTlsCrypto, new TlsPskIdentityManager(PSK_IDENTITY)) { @Override public int[] getCipherSuites() { return CIPHER_SUITES; } }; tlsServerProtocol.accept(pskTlsServer); final InputStream tlsInputStream = tlsServerProtocol.getInputStream(); final OutputStream tlsOutputStream = tlsServerProtocol.getOutputStream(); Streams.pipeAll(tlsInputStream, tlsOutputStream); } else { // PLAINTEXT echo Streams.pipeAll(inputStream, outputStream); } } } } } public Thread start() { final CountDownLatch countDownLatch = new CountDownLatch(1); final Thread serverThread = new Thread(() -> { try { this.listen(countDownLatch); } catch (Exception e) { log.info("Main: Exception while listening", e); } }); log.info("Main: Waiting for Server"); final long nanos = System.nanoTime(); serverThread.start(); try { countDownLatch.await(MAX_WAIT_MILLIS, TimeUnit.MILLISECONDS); // wait for server thread to indicate it started OK // Thread.sleep(100); // waiting for server to call serverSocket.accept() doesn't seem to help } catch (InterruptedException e) { log.info("Main: Exception while waiting for start", e); throw new RuntimeException(e); } finally { log.info("Main: Waited for Server start for " + Float.valueOf((System.nanoTime() - nanos)/1000000F) + " msec"); } return serverThread; } } public static class TlsPskIdentity implements TlsPSKIdentity { private final byte[] identity; private final byte[] psk; public TlsPskIdentity(final byte[] pskIdentity, final byte[] psk) { this.identity = pskIdentity; this.psk = psk; } @Override public byte[] getPSKIdentity() { return this.identity; } @Override public byte[] getPSK() { return this.psk; } @Override public void skipIdentityHint() { /*do nothing*/ } @Override public void notifyIdentityHint(byte[] psk_identity_hint) { /*do nothing*/ } } public static class TlsPskIdentityManager implements TlsPSKIdentityManager { private final TlsPskIdentity tlsPskIdentity; public TlsPskIdentityManager(final TlsPskIdentity tlsPskIdentity) { this.tlsPskIdentity = tlsPskIdentity; } @Override public byte[] getHint() { return this.tlsPskIdentity.getPSKIdentity(); } @Override public byte[] getPSK(byte[] identity) { return this.tlsPskIdentity.getPSK(); } } }
It seems straight forward to reproduce the problem.
As a quick estimate, I have to re-run the JUnit test about 10 times before I randomly get all tests to pass. The rest of the runs, the first TLS PSK test fails, so about a 90% failure rate, but only for the first parameterized test repeat.
The return value from TlsPSKIdentity#getPSK (resp. TlsPSKIdentityManager#getPSK) needs to be cloned, as it will be filled with zeros after use.
You could e.g. use org.bouncycastle.tls.BasicTlsPSKIdentity in this case.
Thank you for your feedback. I tried your suggestion and it worked! My mini stress test passes 100% of the time now. Awesome! I will close this issue.
Q: Is returning a cloned byte[] from TlsPSKIdentity#getPSK documented somewhere?
I hope it is OK to ask this question as a postscript to the original issue.
I don't recall seeing this in main/test comments, or in javadocs. I didn't expect a read method to have a side effect of zeroing out the byte array.
- addedsupport requestCommunity assistance requestedCommunity assistance requested
on Dec 13, 2024
I think I am experiencing a possible race condition in either TlsClientProtocol or TlsServerProtocol.
I am using bctls-jdk18on 1.78.1.
I can reproduce an intermittent failure in a JUnit test. It performs:
Sometimes all 20 tests pass. Other times only 19 out of 20 tests pass.
I uploaded a Maven project to GitHub to demonstrate the issue.
Here are screenshots comparing when all tests passed, versus when all tests passed except the first TLS PSK test.
If TLS PSK test #1 fails, the stack trace is: