diff --git a/src/main/java/org/metricshub/ipmi/core/api/async/IpmiAsyncConnector.java b/src/main/java/org/metricshub/ipmi/core/api/async/IpmiAsyncConnector.java index 3e9c037..15e7d71 100644 --- a/src/main/java/org/metricshub/ipmi/core/api/async/IpmiAsyncConnector.java +++ b/src/main/java/org/metricshub/ipmi/core/api/async/IpmiAsyncConnector.java @@ -44,8 +44,8 @@ import java.io.FileNotFoundException; import java.io.IOException; import java.net.InetAddress; -import java.util.ArrayList; import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; import java.util.concurrent.TimeUnit; import org.slf4j.Logger; @@ -87,8 +87,8 @@ public class IpmiAsyncConnector implements ConnectionListener { private ConnectionManager connectionManager; private SessionManager sessionManager; private int retries; - private final List responseListeners; - private final List inboundMessageListeners; + private final List responseListeners = new CopyOnWriteArrayList(); + private final List inboundMessageListeners = new CopyOnWriteArrayList(); private static Logger logger = LoggerFactory.getLogger(IpmiAsyncConnector.class); @@ -104,8 +104,6 @@ public class IpmiAsyncConnector implements ConnectionListener { * when properties file was not found */ public IpmiAsyncConnector(int port) throws IOException { - responseListeners = new ArrayList(); - inboundMessageListeners = new ArrayList(); connectionManager = new ConnectionManager(port); sessionManager = new SessionManager(); loadProperties(); @@ -126,8 +124,6 @@ public IpmiAsyncConnector(int port) throws IOException { * when properties file was not found */ public IpmiAsyncConnector(int port, InetAddress address) throws IOException { - responseListeners = new ArrayList(); - inboundMessageListeners = new ArrayList(); connectionManager = new ConnectionManager(port, address); sessionManager = new SessionManager(); loadProperties(); @@ -145,8 +141,6 @@ public IpmiAsyncConnector(int port, InetAddress address) throws IOException { * error. */ public IpmiAsyncConnector(int port, long pingPeriod) throws IOException { - responseListeners = new ArrayList<>(); - inboundMessageListeners = new ArrayList<>(); connectionManager = new ConnectionManager(port, pingPeriod); sessionManager = new SessionManager(); loadProperties(); @@ -475,9 +469,7 @@ public int retry(ConnectionHandle connectionHandle, int tag, PayloadType message * {@link IpmiResponseListener} to processResponse */ public void registerListener(IpmiResponseListener listener) { - synchronized (responseListeners) { - responseListeners.add(listener); - } + responseListeners.add(listener); } /** @@ -488,9 +480,7 @@ public void registerListener(IpmiResponseListener listener) { * - the {@link IpmiResponseListener} to unregister */ public void unregisterListener(IpmiResponseListener listener) { - synchronized (responseListeners) { - responseListeners.remove(listener); - } + responseListeners.remove(listener); } /** @@ -500,9 +490,7 @@ public void unregisterListener(IpmiResponseListener listener) { * the {@link InboundMessageListener} to register. */ public void registerIncomingPayloadListener(InboundMessageListener listener) { - synchronized (inboundMessageListeners) { - inboundMessageListeners.add(listener); - } + inboundMessageListeners.add(listener); } /** @@ -512,9 +500,7 @@ public void registerIncomingPayloadListener(InboundMessageListener listener) { * the {@link InboundMessageListener} to unregister. */ public void unregisterIncomingPayloadListener(InboundMessageListener listener) { - synchronized (inboundMessageListeners) { - inboundMessageListeners.remove(listener); - } + inboundMessageListeners.remove(listener); } @Override @@ -539,11 +525,9 @@ public void processResponse(ResponseData responseData, int handle, int tag, Exce new ConnectionHandle(handle, connection.getRemoteMachineAddress(), connection.getRemoteMachinePort())); } - synchronized (responseListeners) { - for (IpmiResponseListener listener : responseListeners) { - if (listener != null) { - listener.notify(response); - } + for (IpmiResponseListener listener : responseListeners) { + if (listener != null) { + listener.notify(response); } } } diff --git a/src/main/java/org/metricshub/ipmi/core/coding/security/ConfidentialityAesCbc128.java b/src/main/java/org/metricshub/ipmi/core/coding/security/ConfidentialityAesCbc128.java index 57cbf77..6f179bb 100644 --- a/src/main/java/org/metricshub/ipmi/core/coding/security/ConfidentialityAesCbc128.java +++ b/src/main/java/org/metricshub/ipmi/core/coding/security/ConfidentialityAesCbc128.java @@ -44,6 +44,7 @@ public class ConfidentialityAesCbc128 extends ConfidentialityAlgorithm { Arrays.fill(CONST2, (byte) 2); } + // The sending and receiving threads share this instance: encrypt() and decrypt() each init the cipher private Cipher cipher; private SecretKeySpec cipherKey; @@ -54,7 +55,7 @@ public byte getCode() { } @Override - public void initialize(byte[] sik, AuthenticationAlgorithm authenticationAlgorithm) + public synchronized void initialize(byte[] sik, AuthenticationAlgorithm authenticationAlgorithm) throws InvalidKeyException, NoSuchAlgorithmException, NoSuchPaddingException { @@ -78,7 +79,7 @@ public void initialize(byte[] sik, AuthenticationAlgorithm authenticationAlgorit } @Override - public byte[] encrypt(byte[] data) throws InvalidKeyException { + public synchronized byte[] encrypt(byte[] data) throws InvalidKeyException { int length = data.length + 17; int pad = 0; if (length % 16 != 0) { @@ -115,7 +116,7 @@ public byte[] encrypt(byte[] data) throws InvalidKeyException { } @Override - public byte[] decrypt(byte[] data) { + public synchronized byte[] decrypt(byte[] data) { byte[] decrypted = null; try { diff --git a/src/main/java/org/metricshub/ipmi/core/coding/security/IntegrityAlgorithm.java b/src/main/java/org/metricshub/ipmi/core/coding/security/IntegrityAlgorithm.java index b739adf..d5684e2 100644 --- a/src/main/java/org/metricshub/ipmi/core/coding/security/IntegrityAlgorithm.java +++ b/src/main/java/org/metricshub/ipmi/core/coding/security/IntegrityAlgorithm.java @@ -67,7 +67,7 @@ private IntegrityAlgorithm(Mac mac) { * @param key - Session Integrity Key calculated during the opening of the * session or user password if 'one-key' logins are enabled. */ - public void initialize(byte[] key) throws InvalidKeyException { + public synchronized void initialize(byte[] key) throws InvalidKeyException { this.sik = key; final String algorithmName = getAlgorithmName(); @@ -103,14 +103,15 @@ protected void setSik(byte[] sik) { public abstract byte getCode(); /** - * Creates AuthCode field for message. + * Creates AuthCode field for message. Synchronized: the sending, receiving and keep-alive threads share this + * instance and its {@link Mac}. * * @param base - data starting with the AuthType/Format field up to and * including the field that immediately precedes the AuthCode field * @return AuthCode field. Might be null if empty AuthCOde field is generated. * @see Rakp1#calculateSik(org.metricshub.ipmi.core.coding.commands.session.Rakp1ResponseData) */ - public byte[] generateAuthCode(final byte[] base) { + public synchronized byte[] generateAuthCode(final byte[] base) { if (sik == null) { throw new NullPointerException("Algorithm not initialized."); diff --git a/src/main/java/org/metricshub/ipmi/core/connection/Connection.java b/src/main/java/org/metricshub/ipmi/core/connection/Connection.java index f36cf12..2491bf2 100644 --- a/src/main/java/org/metricshub/ipmi/core/connection/Connection.java +++ b/src/main/java/org/metricshub/ipmi/core/connection/Connection.java @@ -78,6 +78,7 @@ import java.util.Map; import java.util.Timer; import java.util.TimerTask; +import java.util.concurrent.CopyOnWriteArrayList; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; @@ -140,7 +141,7 @@ public void setTimeout(int timeout) { public Connection(Messenger messenger, int handle) { stateMachine = new StateMachine(messenger); this.handle = handle; - listeners = new ArrayList(); + listeners = new CopyOnWriteArrayList(); timeout = Integer.parseInt(PropertiesManager.getInstance().getProperty("timeout")); messageHandlers = new EnumMap(PayloadType.class); currentSessionSequenceNumber = new AtomicInteger(0); @@ -327,6 +328,7 @@ public List getAvailableCipherSuites(int tag) throws Exception { private void waitForResponse() throws Exception { long deadline = System.nanoTime() + TimeUnit.MILLISECONDS.toNanos(timeout); + InterruptedException interrupted = null; synchronized (responseLock) { try { long remaining = deadline - System.nanoTime(); @@ -335,16 +337,26 @@ private void waitForResponse() throws Exception { remaining = deadline - System.nanoTime(); } } catch (InterruptedException e) { - // The caller gave up on us (Future.cancel): leave the state machine in a state that allows a retry - stateMachine.doTransition(new Timeout()); - Thread.currentThread().interrupt(); - throw e; + interrupted = e; } } + if (interrupted != null) { + // The caller gave up on us (Future.cancel): leave the state machine in a state that allows a retry. + // Outside the response lock: the receiving thread takes it while holding the state machine lock. + stateMachine.doTransition(new Timeout()); + Thread.currentThread().interrupt(); + throw interrupted; + } if (lastAction == null) { - stateMachine.doTransition(new Timeout()); - throw new ConnectionException("Command timed out"); + // The receiving thread publishes a reply under the state machine lock: once we hold it, a reply that + // is not there yet cannot be processed before the timeout rolls the state back + synchronized (stateMachine) { + if (lastAction == null) { + stateMachine.doTransition(new Timeout()); + throw new ConnectionException("Command timed out"); + } + } } if (!(lastAction instanceof ResponseAction || lastAction instanceof GetSikAction)) { if (lastAction instanceof ErrorAction) { diff --git a/src/main/java/org/metricshub/ipmi/core/sm/StateMachine.java b/src/main/java/org/metricshub/ipmi/core/sm/StateMachine.java index a9890bb..babbffb 100644 --- a/src/main/java/org/metricshub/ipmi/core/sm/StateMachine.java +++ b/src/main/java/org/metricshub/ipmi/core/sm/StateMachine.java @@ -26,9 +26,11 @@ import java.net.InetAddress; import java.util.ArrayList; import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; import org.metricshub.ipmi.core.coding.rmcp.RmcpDecoder; import org.metricshub.ipmi.core.common.Constants; +import org.metricshub.ipmi.core.sm.actions.MessageAction; import org.metricshub.ipmi.core.sm.actions.StateMachineAction; import org.metricshub.ipmi.core.sm.events.StateMachineEvent; import org.metricshub.ipmi.core.sm.states.SessionValid; @@ -44,15 +46,23 @@ */ public class StateMachine implements UdpListener { - private List observers; + private final List observers = new CopyOnWriteArrayList(); - private State current; + /** + * In-session messages received while the lock is held, dispatched by {@link #doTransition(StateMachineEvent)} + * and {@link #notifyMessage(UdpMessage)} once they release it, so that no application listener runs under the + * lock. The other actions (a handshake reply, an error, the session key) are published under the lock: the + * caller then sees the reply together with the state it produced, and cannot time the request out in between. + */ + private final List pendingActions = new ArrayList(); + + private volatile State current; private Messenger messenger; - private InetAddress remoteMachineAddress; - private int remoteMachinePort; + private volatile InetAddress remoteMachineAddress; + private volatile int remoteMachinePort; - private boolean initialized; + private volatile boolean initialized; public State getCurrent() { return current; @@ -72,7 +82,6 @@ public void setCurrent(State current) { */ public StateMachine(Messenger messenger) { this.messenger = messenger; - observers = new ArrayList(); initialized = false; } @@ -101,12 +110,21 @@ public int getRemoteMachinePort() { } /** - * Sends a notification of an action to all {@link MachineObserver}s + * Sends a notification of an action to all {@link MachineObserver}s. A {@link MessageAction} emitted by a state + * while the lock is held is deferred until the transition releases it; the other actions are published at once. * * @param action * - a {@link StateMachineAction} to perform */ public void doExternalAction(StateMachineAction action) { + if (action instanceof MessageAction && Thread.holdsLock(this)) { + pendingActions.add(action); + } else { + notifyObservers(action); + } + } + + private void notifyObservers(StateMachineAction action) { for (MachineObserver observer : observers) { if (observer != null) { observer.notify(action); @@ -114,6 +132,19 @@ public void doExternalAction(StateMachineAction action) { } } + private void dispatch(List actions) { + for (StateMachineAction action : actions) { + notifyObservers(action); + } + } + + /** Returns the actions emitted so far and clears them; called under the lock. */ + private List drainPendingActions() { + List actions = new ArrayList(pendingActions); + pendingActions.clear(); + return actions; + } + /** * Sets the State Machine in the initial state. * @@ -154,7 +185,9 @@ public boolean isActive() { /** * Performs a {@link State} transition according to the event and - * {@link #current} state + * {@link #current} state. Transitions and received messages are serialized, so a late reply cannot interleave + * with the timeout or close of the request it answers; the in-session messages are dispatched once the lock is + * released. * * @param event * - {@link StateMachineEvent} invoking the transition @@ -163,17 +196,28 @@ public boolean isActive() { * @see #start(InetAddress, int) */ public void doTransition(StateMachineEvent event) { - if (!initialized) { - throw new NullPointerException("State machine not started"); + List actions; + synchronized (this) { + if (!initialized) { + throw new NullPointerException("State machine not started"); + } + current.doTransition(this, event); + actions = drainPendingActions(); } - current.doTransition(this, event); + dispatch(actions); } @Override public void notifyMessage(UdpMessage message) { - if (message.getAddress().equals(getRemoteMachineAddress()) && message.getPort() == getRemoteMachinePort()) { + if (!message.getAddress().equals(getRemoteMachineAddress()) || message.getPort() != getRemoteMachinePort()) { + return; + } + List actions; + synchronized (this) { current.doAction(this, RmcpDecoder.decode(message.getMessage())); + actions = drainPendingActions(); } + dispatch(actions); } /** diff --git a/src/site/markdown/upgrading.md b/src/site/markdown/upgrading.md index d398f5d..b47d005 100644 --- a/src/site/markdown/upgrading.md +++ b/src/site/markdown/upgrading.md @@ -48,7 +48,11 @@ The `IpmiClient` API is unchanged, and the client is more tolerant of real-world whole connector; * the `PropertiesManager` lookups are logged at `DEBUG` instead of `INFO`, the unused `cleaningFrequency` property is gone from `connection.properties`, and `Constants.TIMEOUT`, - which nothing reads, is deprecated. + which nothing reads, is deprecated; +* the sending, receiving and keep-alive threads of a connection no longer race: the state + machine serializes transitions and received messages, the HMAC and AES objects of a cipher suite + are used by one thread at a time, and the listener lists can be changed while they are being + notified (a listener may unregister itself from its own callback). The decoders follow the IPMI 2.0 and FRU specifications more closely; the visible changes are: diff --git a/src/test/java/org/metricshub/ipmi/core/api/async/IpmiAsyncConnectorTest.java b/src/test/java/org/metricshub/ipmi/core/api/async/IpmiAsyncConnectorTest.java new file mode 100644 index 0000000..16dbfa5 --- /dev/null +++ b/src/test/java/org/metricshub/ipmi/core/api/async/IpmiAsyncConnectorTest.java @@ -0,0 +1,64 @@ +package org.metricshub.ipmi.core.api.async; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.net.InetAddress; +import java.util.concurrent.atomic.AtomicInteger; + +import org.junit.jupiter.api.Test; +import org.metricshub.ipmi.core.api.async.messages.IpmiResponse; +import org.metricshub.ipmi.core.coding.payload.IpmiPayload; +import org.metricshub.ipmi.core.coding.payload.PlainMessage; + +class IpmiAsyncConnectorTest { + + @Test + void aResponseListenerMayUnregisterItselfWhileNotified() throws Exception { + IpmiAsyncConnector connector = new IpmiAsyncConnector(0); + try { + ConnectionHandle handle = connector.createConnection(InetAddress.getLoopbackAddress(), 623); + AtomicInteger notified = new AtomicInteger(); + IpmiResponseListener oneShot = new IpmiResponseListener() { + @Override + public void notify(IpmiResponse response) { + connector.unregisterListener(this); + } + }; + connector.registerListener(oneShot); + connector.registerListener(response -> notified.incrementAndGet()); + + connector.processResponse(null, handle.getHandle(), 1, new Exception("timed out")); + connector.processResponse(null, handle.getHandle(), 2, new Exception("timed out")); + assertEquals(2, notified.get(), "the listener registered after the one-shot one must be notified each time"); + } finally { + connector.tearDown(); + } + } + + @Test + void anInboundListenerMayUnregisterItselfWhileNotified() throws Exception { + IpmiAsyncConnector connector = new IpmiAsyncConnector(0); + try { + AtomicInteger notified = new AtomicInteger(); + InboundMessageListener oneShot = new InboundMessageListener() { + @Override + public boolean isPayloadSupported(IpmiPayload payload) { + return true; + } + + @Override + public void notify(IpmiPayload payload) { + notified.incrementAndGet(); + connector.unregisterIncomingPayloadListener(this); + } + }; + connector.registerIncomingPayloadListener(oneShot); + connector.processRequest(new PlainMessage(new byte[0])); + connector.processRequest(new PlainMessage(new byte[0])); + assertTrue(notified.get() == 1, "the one-shot listener was notified " + notified.get() + " times"); + } finally { + connector.tearDown(); + } + } +} diff --git a/src/test/java/org/metricshub/ipmi/core/coding/security/ConfidentialityAesCbc128Test.java b/src/test/java/org/metricshub/ipmi/core/coding/security/ConfidentialityAesCbc128Test.java new file mode 100644 index 0000000..36f2023 --- /dev/null +++ b/src/test/java/org/metricshub/ipmi/core/coding/security/ConfidentialityAesCbc128Test.java @@ -0,0 +1,44 @@ +package org.metricshub.ipmi.core.coding.security; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertNull; + +import java.util.Arrays; +import java.util.concurrent.atomic.AtomicReference; + +import org.junit.jupiter.api.Test; + +class ConfidentialityAesCbc128Test { + + private static final int THREADS = 8; + + private static final int ROUNDS = 2000; + + @Test + void encryptAndDecryptRoundTripWhenSeveralThreadsShareTheAlgorithm() throws Exception { + ConfidentialityAesCbc128 algorithm = new ConfidentialityAesCbc128(); + algorithm.initialize(new byte[20], new AuthenticationRakpHmacSha1()); + + AtomicReference failure = new AtomicReference<>(); + Thread[] threads = new Thread[THREADS]; + for (int i = 0; i < THREADS; i++) { + final byte[] data = new byte[20 + i]; + Arrays.fill(data, (byte) (i + 1)); + threads[i] = new Thread(() -> { + for (int round = 0; round < ROUNDS; round++) { + try { + assertArrayEquals(data, algorithm.decrypt(algorithm.encrypt(data))); + } catch (Throwable e) { + failure.compareAndSet(null, e); + return; + } + } + }); + threads[i].start(); + } + for (Thread thread : threads) { + thread.join(); + } + assertNull(failure.get(), String.valueOf(failure.get())); + } +} diff --git a/src/test/java/org/metricshub/ipmi/core/coding/security/IntegrityAlgorithmTest.java b/src/test/java/org/metricshub/ipmi/core/coding/security/IntegrityAlgorithmTest.java new file mode 100644 index 0000000..667bbae --- /dev/null +++ b/src/test/java/org/metricshub/ipmi/core/coding/security/IntegrityAlgorithmTest.java @@ -0,0 +1,49 @@ +package org.metricshub.ipmi.core.coding.security; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertNull; + +import java.util.Arrays; +import java.util.concurrent.atomic.AtomicReference; + +import org.junit.jupiter.api.Test; + +class IntegrityAlgorithmTest { + + private static final int THREADS = 8; + + private static final int ROUNDS = 2000; + + @Test + void generateAuthCodeIsStableWhenSeveralThreadsShareTheAlgorithm() throws Exception { + IntegrityAlgorithm algorithm = new IntegrityHmacSha1_96(); + algorithm.initialize(new byte[20]); + byte[][] bases = new byte[THREADS][40]; + byte[][] expected = new byte[THREADS][]; + for (int i = 0; i < THREADS; i++) { + Arrays.fill(bases[i], (byte) (i + 1)); + expected[i] = algorithm.generateAuthCode(bases[i]); + } + + AtomicReference failure = new AtomicReference<>(); + Thread[] threads = new Thread[THREADS]; + for (int i = 0; i < THREADS; i++) { + final int id = i; + threads[i] = new Thread(() -> { + for (int round = 0; round < ROUNDS; round++) { + try { + assertArrayEquals(expected[id], algorithm.generateAuthCode(bases[id]), "thread " + id); + } catch (AssertionError e) { + failure.compareAndSet(null, e); + return; + } + } + }); + threads[i].start(); + } + for (Thread thread : threads) { + thread.join(); + } + assertNull(failure.get(), String.valueOf(failure.get())); + } +} diff --git a/src/test/java/org/metricshub/ipmi/core/connection/ConnectionTest.java b/src/test/java/org/metricshub/ipmi/core/connection/ConnectionTest.java index 235a55b..ef05066 100644 --- a/src/test/java/org/metricshub/ipmi/core/connection/ConnectionTest.java +++ b/src/test/java/org/metricshub/ipmi/core/connection/ConnectionTest.java @@ -32,6 +32,10 @@ import org.metricshub.ipmi.core.sm.actions.MessageAction; import org.metricshub.ipmi.core.sm.states.SessionValid; import org.metricshub.ipmi.core.transport.UdpMessage; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import org.metricshub.ipmi.core.coding.commands.session.GetChannelCipherSuitesResponseData; +import org.metricshub.ipmi.core.sm.actions.ResponseAction; class ConnectionTest { @@ -330,4 +334,78 @@ public void processRequest(IpmiPayload payload) { connection.disconnect(); } } + + @Test + void aListenerMayUnregisterItselfWhileNotified() throws Exception { + Connection connection = connect(TIMEOUT_MS); + try { + AtomicInteger notified = new AtomicInteger(); + ConnectionListener oneShot = new ConnectionListener() { + @Override + public void processResponse(ResponseData responseData, int handle, int tag, Exception exception) { + connection.unregisterListener(this); + } + + @Override + public void processRequest(IpmiPayload payload) { + connection.unregisterListener(this); + } + }; + connection.registerListener(oneShot); + connection.registerListener(new ConnectionListener() { + @Override + public void processResponse(ResponseData responseData, int handle, int tag, Exception exception) { + notified.incrementAndGet(); + } + + @Override + public void processRequest(IpmiPayload payload) { + notified.incrementAndGet(); + } + }); + connection.notifyResponseListeners(0, 1, null, new Exception("timed out")); + connection.notifyResponseListeners(0, 2, null, new Exception("timed out")); + assertEquals(2, notified.get(), "the second listener must be notified each time"); + } finally { + connection.disconnect(); + } + } + + @Test + void aReplyPublishedWhileTheTimeoutIsPendingIsNotLost() throws Exception { + Connection connection = connect(TIMEOUT_MS); + try { + Field field = Connection.class.getDeclaredField("stateMachine"); + field.setAccessible(true); + StateMachine machine = (StateMachine) field.get(connection); + AtomicReference publisherFailure = new AtomicReference<>(); + CountDownLatch published = new CountDownLatch(1); + // Plays the receiving thread: takes the state machine lock once the request is sent, keeps it past the + // deadline (the waiter times out meanwhile and blocks on the lock), then publishes the reply under it + Thread publisher = new Thread(() -> { + try { + Thread.sleep(TIMEOUT_MS / 2); + synchronized (machine) { + Thread.sleep(2 * TIMEOUT_MS); + GetChannelCipherSuitesResponseData data = new GetChannelCipherSuitesResponseData(); + data.setCipherSuiteData(new byte[0]); + connection.notify(new ResponseAction(data)); + published.countDown(); + } + } catch (Throwable t) { + publisherFailure.set(t); + } + }); + publisher.start(); + + List suites = connection.getAvailableCipherSuites(1); + + publisher.join(5000); + assertEquals(null, publisherFailure.get()); + assertEquals(0, published.getCount(), "the reply was published before the waiter could proceed"); + assertTrue(suites.isEmpty(), "the reply, not a timeout, ends the step"); + } finally { + connection.disconnect(); + } + } } diff --git a/src/test/java/org/metricshub/ipmi/core/sm/StateMachineTest.java b/src/test/java/org/metricshub/ipmi/core/sm/StateMachineTest.java new file mode 100644 index 0000000..6b577fb --- /dev/null +++ b/src/test/java/org/metricshub/ipmi/core/sm/StateMachineTest.java @@ -0,0 +1,88 @@ +package org.metricshub.ipmi.core.sm; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.net.InetAddress; +import java.util.List; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.concurrent.atomic.AtomicBoolean; + +import org.junit.jupiter.api.Test; +import org.metricshub.ipmi.core.coding.Encoder; +import org.metricshub.ipmi.core.coding.commands.IpmiVersion; +import org.metricshub.ipmi.core.coding.commands.session.GetChannelAuthenticationCapabilities; +import org.metricshub.ipmi.core.coding.protocol.encoder.Protocolv20Encoder; +import org.metricshub.ipmi.core.coding.security.CipherSuite; +import org.metricshub.ipmi.core.sm.actions.ErrorAction; +import org.metricshub.ipmi.core.sm.actions.MessageAction; +import org.metricshub.ipmi.core.sm.actions.StateMachineAction; +import org.metricshub.ipmi.core.sm.events.GetChannelCipherSuitesPending; +import org.metricshub.ipmi.core.sm.states.SessionValid; +import org.metricshub.ipmi.core.sm.states.Uninitialized; +import org.metricshub.ipmi.core.transport.SilentMessenger; +import org.metricshub.ipmi.core.transport.UdpMessage; + +class StateMachineTest { + + private static final int SESSION_ID = 1; + + private final List actions = new CopyOnWriteArrayList<>(); + + private final AtomicBoolean lockHeld = new AtomicBoolean(); + + private final AtomicBoolean rolledBack = new AtomicBoolean(); + + private StateMachine machine; + + private void start(SilentMessenger messenger) { + machine = new StateMachine(messenger); + machine.start(InetAddress.getLoopbackAddress(), 623); + machine.register(action -> { + actions.add(action); + lockHeld.set(Thread.holdsLock(machine)); + rolledBack.set(machine.getCurrent() instanceof Uninitialized); + }); + } + + @Test + void aHandshakeErrorIsPublishedWithTheStateItProduced() { + start(new SilentMessenger() { + @Override + public void send(UdpMessage message) { + throw new IllegalStateException("cable unplugged"); + } + }); + + machine.doTransition(new GetChannelCipherSuitesPending(1)); + + assertEquals(1, actions.size()); + assertTrue(actions.get(0) instanceof ErrorAction, String.valueOf(actions.get(0))); + assertTrue(rolledBack.get(), "the state must be rolled back when the error is published"); + assertTrue(lockHeld.get(), "a handshake action is published under the lock, before a timeout can be applied"); + } + + @Test + void anInSessionMessageIsDispatchedOutsideTheLock() throws Exception { + start(new SilentMessenger()); + machine.setCurrent(new SessionValid(CipherSuite.getEmpty(), SESSION_ID)); + byte[] raw = Encoder + .encode( + new Protocolv20Encoder(), + new GetChannelAuthenticationCapabilities(IpmiVersion.V20, IpmiVersion.V20, CipherSuite.getEmpty()), + 1, + 1, + SESSION_ID); + UdpMessage message = new UdpMessage(); + message.setAddress(InetAddress.getLoopbackAddress()); + message.setPort(623); + message.setMessage(raw); + + machine.notifyMessage(message); + + assertEquals(1, actions.size()); + assertTrue(actions.get(0) instanceof MessageAction, String.valueOf(actions.get(0))); + assertFalse(lockHeld.get(), "an in-session message must not be dispatched under the lock"); + } +}