Skip to content

TLS PSK intermittently fails with TlsFatalAlertReceived: bad_record_mac(20) #1876

Description

@justincranford

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:

  • 10 repeats of a client/server echo test, using plaintext communication
  • 10 repeats of a client/server echo test, using TLS PSK communication

Sometimes all 20 tests pass. Other times only 19 out of 20 tests pass.

  1. Plaintext tests 1 through 10 always pass.
  2. TLS PSK test 1 intermittently passes, or fails with exception TlsFatalAlertReceived: bad_record_mac(20).
  3. TLS PSK tests 2 through 10 always 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.

image
image

If TLS PSK test #1 fails, the stack trace is:

com.github.justincranford.psk.PskTlsTest
testTlsPsk(com.github.justincranford.psk.PskTlsTest)
org.bouncycastle.tls.TlsFatalAlertReceived: bad_record_mac(20)
	at org.bouncycastle.tls.TlsProtocol.handleAlertMessage(TlsProtocol.java:245)
	at org.bouncycastle.tls.TlsProtocol.processAlertQueue(TlsProtocol.java:740)
	at org.bouncycastle.tls.TlsProtocol.processRecord(TlsProtocol.java:563)
	at org.bouncycastle.tls.RecordStream.readRecord(RecordStream.java:247)
	at org.bouncycastle.tls.TlsProtocol.safeReadRecord(TlsProtocol.java:879)
	at org.bouncycastle.tls.TlsProtocol.blockForHandshake(TlsProtocol.java:427)
	at org.bouncycastle.tls.TlsClientProtocol.connect(TlsClientProtocol.java:88)
	at com.github.justincranford.psk.PskTlsTest$PskTlsClient.send(PskTlsTest.java:79)
	at com.github.justincranford.psk.PskTlsTest.doClientServer(PskTlsTest.java:62)
	at com.github.justincranford.psk.PskTlsTest.testTlsPsk(PskTlsTest.java:51)
	at java.base/java.lang.reflect.Method.invoke(Method.java:580)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.accept(ForEachOps.java:184)
	at java.base/java.util.stream.ReferencePipeline$3$1.accept(ReferencePipeline.java:197)
	at java.base/java.util.stream.ReferencePipeline$2$1.accept(ReferencePipeline.java:179)
	at java.base/java.util.stream.ReferencePipeline$3$1.accept(ReferencePipeline.java:197)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.accept(ForEachOps.java:184)
	at java.base/java.util.stream.ReferencePipeline$3$1.accept(ReferencePipeline.java:197)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.accept(ForEachOps.java:184)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.accept(ForEachOps.java:184)
	at java.base/java.util.stream.ReferencePipeline$3$1.accept(ReferencePipeline.java:197)
	at java.base/java.util.Spliterators$ArraySpliterator.forEachRemaining(Spliterators.java:1024)
	at java.base/java.util.stream.AbstractPipeline.copyInto(AbstractPipeline.java:509)
	at java.base/java.util.stream.AbstractPipeline.wrapAndCopyInto(AbstractPipeline.java:499)
	at java.base/java.util.stream.ForEachOps$ForEachOp.evaluateSequential(ForEachOps.java:151)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.evaluateSequential(ForEachOps.java:174)
	at java.base/java.util.stream.AbstractPipeline.evaluate(AbstractPipeline.java:234)
	at java.base/java.util.stream.ReferencePipeline.forEach(ReferencePipeline.java:596)
	at java.base/java.util.stream.ReferencePipeline$7$1.accept(ReferencePipeline.java:276)
	at java.base/java.util.ArrayList$ArrayListSpliterator.forEachRemaining(ArrayList.java:1708)
	at java.base/java.util.stream.AbstractPipeline.copyInto(AbstractPipeline.java:509)
	at java.base/java.util.stream.AbstractPipeline.wrapAndCopyInto(AbstractPipeline.java:499)
	at java.base/java.util.stream.ForEachOps$ForEachOp.evaluateSequential(ForEachOps.java:151)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.evaluateSequential(ForEachOps.java:174)
	at java.base/java.util.stream.AbstractPipeline.evaluate(AbstractPipeline.java:234)
	at java.base/java.util.stream.ReferencePipeline.forEach(ReferencePipeline.java:596)
	at java.base/java.util.stream.ReferencePipeline$7$1.accept(ReferencePipeline.java:276)
	at java.base/java.util.stream.ReferencePipeline$3$1.accept(ReferencePipeline.java:197)
	at java.base/java.util.stream.ReferencePipeline$3$1.accept(ReferencePipeline.java:197)
	at java.base/java.util.stream.ReferencePipeline$3$1.accept(ReferencePipeline.java:197)
	at java.base/java.util.ArrayList$ArrayListSpliterator.forEachRemaining(ArrayList.java:1708)
	at java.base/java.util.stream.AbstractPipeline.copyInto(AbstractPipeline.java:509)
	at java.base/java.util.stream.AbstractPipeline.wrapAndCopyInto(AbstractPipeline.java:499)
	at java.base/java.util.stream.ForEachOps$ForEachOp.evaluateSequential(ForEachOps.java:151)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.evaluateSequential(ForEachOps.java:174)
	at java.base/java.util.stream.AbstractPipeline.evaluate(AbstractPipeline.java:234)
	at java.base/java.util.stream.ReferencePipeline.forEach(ReferencePipeline.java:596)
	at java.base/java.util.stream.ReferencePipeline$7$1.accept(ReferencePipeline.java:276)
	at java.base/java.util.ArrayList$ArrayListSpliterator.forEachRemaining(ArrayList.java:1708)
	at java.base/java.util.stream.AbstractPipeline.copyInto(AbstractPipeline.java:509)
	at java.base/java.util.stream.AbstractPipeline.wrapAndCopyInto(AbstractPipeline.java:499)
	at java.base/java.util.stream.ForEachOps$ForEachOp.evaluateSequential(ForEachOps.java:151)
	at java.base/java.util.stream.ForEachOps$ForEachOp$OfRef.evaluateSequential(ForEachOps.java:174)
	at java.base/java.util.stream.AbstractPipeline.evaluate(AbstractPipeline.java:234)
	at java.base/java.util.stream.ReferencePipeline.forEach(ReferencePipeline.java:596)
	at java.base/java.util.ArrayList.forEach(ArrayList.java:1596)
	at java.base/java.util.ArrayList.forEach(ArrayList.java:1596)

Activity

  1. justincranford commented on Oct 23, 2024

    @justincranford
    Author
    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(); }
    	}
    }
  2. justincranford commented on Oct 27, 2024

    @justincranford
    Author

    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.

  3. peterdettman commented on Oct 28, 2024

    @peterdettman
    Collaborator

    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.

  4. justincranford commented on Oct 30, 2024

    @justincranford
    Author

    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.

  5. justincranford commented on Oct 30, 2024

    @justincranford
    Author

    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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    support requestCommunity assistance requested

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions