Browse Source

Fixes after review

pull/15301/head
Andrii Landiak 6 months ago
parent
commit
1259fdbced
  1. 40
      common/transport/lwm2m/src/main/java/org/thingsboard/server/transport/lwm2m/server/DefaultLwM2mTransportService.java
  2. 2
      common/transport/lwm2m/src/test/java/org/thingsboard/server/transport/lwm2m/config/LwM2MTransportServerConfigDebounceTest.java
  3. 12
      common/transport/mqtt/src/main/java/org/thingsboard/server/transport/mqtt/MqttSslHandlerProvider.java
  4. 33
      common/transport/mqtt/src/test/java/org/thingsboard/server/transport/mqtt/MqttSslHandlerProviderTest.java
  5. 8
      common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/config/ssl/KeystoreSslCredentials.java
  6. 8
      common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/config/ssl/PemSslCredentials.java
  7. 23
      common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/config/ssl/SslCredentialsConfig.java
  8. 37
      common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/service/CertificateReloadManager.java
  9. 7
      common/transport/transport-api/src/test/java/org/thingsboard/server/common/transport/config/ssl/SslCredentialsConfigTest.java

40
common/transport/lwm2m/src/main/java/org/thingsboard/server/transport/lwm2m/server/DefaultLwM2mTransportService.java

@ -32,6 +32,7 @@ import org.eclipse.leshan.server.californium.LwM2mPskStore;
import org.eclipse.leshan.server.californium.endpoint.CaliforniumServerEndpointsProvider;
import org.eclipse.leshan.server.californium.endpoint.coap.CoapServerProtocolProvider;
import org.eclipse.leshan.server.californium.endpoint.coaps.CoapsServerProtocolProvider;
import org.eclipse.leshan.server.endpoint.LwM2mServerEndpointsProvider;
import org.eclipse.leshan.server.registration.RegistrationStore;
import org.springframework.beans.factory.SmartInitializingSingleton;
import org.springframework.context.annotation.DependsOn;
@ -236,36 +237,41 @@ public class DefaultLwM2mTransportService implements LwM2MTransportService, Smar
log.info("Creating new LwM2M server with updated certificates...");
LeshanServer newServer = getLhServer();
// Stop (not destroy) old server to release ports but keep it restartable for rollback
// Only cycle the endpoint providers (CoAP/DTLS). The RegistrationStore and SecurityStore are
// Spring singletons shared with newServer — calling oldServer.stop()/destroy() would propagate
// to them (LeshanServer.stop/destroy propagate to Stoppable/Destroyable stores), which would
// shut down the shared schedulers (TbInMemoryRegistrationStore.destroy calls schedExecutor.shutdownNow),
// killing newServer's cleaner tasks. Leaving the stores running preserves existing device
// registrations across the swap — clients only need to re-establish DTLS on next uplink.
if (oldServer != null) {
log.info("Stopping old LwM2M server to release ports...");
log.info("Stopping old LwM2M endpoints to release ports...");
if (oldListener != null) {
oldServer.getRegistrationService().removeListener(oldListener.registrationListener);
oldServer.getPresenceService().removeListener(oldListener.presenceListener);
oldServer.getObservationService().removeListener(oldListener.observationListener);
oldServer.getSendService().removeListener(oldListener.sendListener);
}
oldServer.stop();
stopEndpoints(oldServer);
}
try {
newServer.start();
} catch (Exception e) {
log.error("Failed to start new LwM2M server", e);
newServer.destroy();
// Attempt to restart the old server (only stopped, not destroyed)
destroyEndpoints(newServer);
// Attempt to restart the old endpoints (shared stores are still running).
if (oldServer != null) {
try {
oldServer.start();
startEndpoints(oldServer);
if (oldListener != null) {
oldServer.getRegistrationService().addListener(oldListener.registrationListener);
oldServer.getPresenceService().addListener(oldListener.presenceListener);
oldServer.getObservationService().addListener(oldListener.observationListener);
oldServer.getSendService().addListener(oldListener.sendListener);
}
log.info("Restored old LwM2M server successfully.");
log.info("Restored old LwM2M endpoints successfully.");
} catch (Exception restoreEx) {
log.error("Failed to restore old LwM2M server", restoreEx);
log.error("Failed to restore old LwM2M endpoints", restoreEx);
}
}
throw e;
@ -280,14 +286,26 @@ public class DefaultLwM2mTransportService implements LwM2MTransportService, Smar
this.server = newServer;
this.context.setServer(newServer);
this.serverListener = newListener;
log.info("New LwM2M server started successfully.");
log.info("New LwM2M server started with refreshed certificates. Existing device registrations preserved; clients will re-establish DTLS on next uplink.");
// Destroy old server only after successful swap
// Destroy old endpoints only — leave the shared stores alone.
if (oldServer != null) {
oldServer.destroy();
destroyEndpoints(oldServer);
}
}
private void stopEndpoints(LeshanServer server) {
server.getEndpointsProvider().forEach(LwM2mServerEndpointsProvider::stop);
}
private void startEndpoints(LeshanServer server) {
server.getEndpointsProvider().forEach(LwM2mServerEndpointsProvider::start);
}
private void destroyEndpoints(LeshanServer server) {
server.getEndpointsProvider().forEach(LwM2mServerEndpointsProvider::destroy);
}
@Override
public String getName() {
return DataConstants.LWM2M_TRANSPORT_NAME;

2
common/transport/lwm2m/src/test/java/org/thingsboard/server/transport/lwm2m/config/LwM2MTransportServerConfigDebounceTest.java

@ -33,7 +33,7 @@ import static org.awaitility.Awaitility.await;
@ExtendWith(MockitoExtension.class)
public class LwM2MTransportServerConfigDebounceTest {
private static final long DEBOUNCE_SECONDS = 2; // matches LwM2MTransportServerConfig.RELOAD_DEBOUNCE_SECONDS
private static final long DEBOUNCE_SECONDS = (long) ReflectionTestUtils.getField(LwM2MTransportServerConfig.class, "RELOAD_DEBOUNCE_SECONDS");
@Mock
private SslCredentialsConfig credentialsConfig;

12
common/transport/mqtt/src/main/java/org/thingsboard/server/transport/mqtt/MqttSslHandlerProvider.java

@ -71,15 +71,21 @@ public class MqttSslHandlerProvider implements SmartInitializingSingleton {
@Override
public void afterSingletonsInstantiated() {
// Eagerly build the initial context so the handshake path is a lock-free volatile read.
this.sslContext = createSslContext();
mqttSslCredentialsConfig.registerReloadCallback(() -> {
log.info("MQTT SSL certificates reloaded. Invalidating SSL context...");
sslContext = null;
log.info("MQTT SSL context invalidated. Will be recreated on next connection.");
log.info("MQTT SSL certificates reloaded. Rebuilding SSL context...");
// Build the new context first; if it fails, the old one stays in place, and
// the exception propagates to CertificateReloadManager's retry/backoff logic.
this.sslContext = createSslContext();
log.info("MQTT SSL context rebuilt. New connections will use the new certificate.");
});
}
public SslHandler getSslHandler() {
SSLContext ctx = sslContext;
// Defensive lazy init in case afterSingletonsInstantiated hasn't run yet (e.g., test wiring).
// In normal operation ctx is non-null here, so the handshake path is lock-free.
if (ctx == null) {
synchronized (this) {
ctx = sslContext;

33
common/transport/mqtt/src/test/java/org/thingsboard/server/transport/mqtt/MqttSslHandlerProviderTest.java

@ -89,38 +89,33 @@ public class MqttSslHandlerProviderTest {
}
@Test
public void givenCertificatesReloaded_whenGetSslHandler_thenShouldRecreateSSLContext() {
public void givenCertificatesReloaded_whenReloadCallbackInvoked_thenShouldRebuildSSLContextEagerly() {
sslHandlerProvider.afterSingletonsInstantiated();
ArgumentCaptor<Runnable> callbackCaptor = ArgumentCaptor.forClass(Runnable.class);
verify(mockCredentialsConfig).registerReloadCallback(callbackCaptor.capture());
Runnable reloadCallback = callbackCaptor.getValue();
SslHandler handler1 = sslHandlerProvider.getSslHandler();
SSLContext initialContext = (SSLContext) ReflectionTestUtils.getField(sslHandlerProvider, "sslContext");
assertThat(initialContext).isNotNull();
reloadCallback.run();
assertThat(handler1).isNotNull();
// After reload the context is rebuilt eagerly (no null-invalidation), so handshakes stay lock-free.
SSLContext contextAfterReload = (SSLContext) ReflectionTestUtils.getField(sslHandlerProvider, "sslContext");
assertThat(contextAfterReload).isNull();
assertThat(contextAfterReload).isNotNull();
assertThat(contextAfterReload).isNotSameAs(initialContext);
SslHandler handler2 = sslHandlerProvider.getSslHandler();
SSLContext newContext = (SSLContext) ReflectionTestUtils.getField(sslHandlerProvider, "sslContext");
assertThat(handler2).isNotNull();
assertThat(newContext).isNotNull();
assertThat(newContext).isNotSameAs(initialContext);
SslHandler handler = sslHandlerProvider.getSslHandler();
assertThat(handler).isNotNull();
}
@Test
public void givenConcurrentGetSslHandlerCalls_whenSSLContextNull_thenShouldCreateOnlyOnce() throws Exception {
public void givenConcurrentGetSslHandlerCalls_whenContextAlreadyBuilt_thenAllReadsReturnSameContext() throws Exception {
sslHandlerProvider.afterSingletonsInstantiated();
ArgumentCaptor<Runnable> callbackCaptor = ArgumentCaptor.forClass(Runnable.class);
verify(mockCredentialsConfig).registerReloadCallback(callbackCaptor.capture());
callbackCaptor.getValue().run();
SSLContext contextBefore = (SSLContext) ReflectionTestUtils.getField(sslHandlerProvider, "sslContext");
assertThat(contextBefore).isNotNull();
CountDownLatch startLatch = new CountDownLatch(1);
CountDownLatch doneLatch = new CountDownLatch(5);
@ -143,12 +138,13 @@ public class MqttSslHandlerProviderTest {
boolean completed = doneLatch.await(5, TimeUnit.SECONDS);
assertThat(completed).isTrue();
SSLContext context = (SSLContext) ReflectionTestUtils.getField(sslHandlerProvider, "sslContext");
assertThat(context).isNotNull();
// Concurrent handshakes read the same pre-built context without the old sync bottleneck.
SSLContext contextAfter = (SSLContext) ReflectionTestUtils.getField(sslHandlerProvider, "sslContext");
assertThat(contextAfter).isSameAs(contextBefore);
}
@Test
public void givenReloadCallback_whenInvoked_thenShouldInvalidateSSLContext() {
public void givenReloadCallback_whenInvoked_thenShouldSwapSSLContextEagerly() {
sslHandlerProvider.afterSingletonsInstantiated();
sslHandlerProvider.getSslHandler();
@ -161,7 +157,8 @@ public class MqttSslHandlerProviderTest {
callbackCaptor.getValue().run();
SSLContext contextAfterReload = (SSLContext) ReflectionTestUtils.getField(sslHandlerProvider, "sslContext");
assertThat(contextAfterReload).isNull();
assertThat(contextAfterReload).isNotNull();
assertThat(contextAfterReload).isNotSameAs(initialContext);
}
@Test

8
common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/config/ssl/KeystoreSslCredentials.java

@ -22,7 +22,6 @@ import org.thingsboard.server.common.data.StringUtils;
import java.io.IOException;
import java.io.InputStream;
import java.nio.file.Files;
import java.nio.file.Path;
import java.security.GeneralSecurityException;
import java.security.KeyStore;
@ -62,10 +61,9 @@ public class KeystoreSslCredentials extends AbstractSslCredentials {
@Override
public List<Path> getCertificateFilePaths() {
if (!StringUtils.isEmpty(storeFile) && !storeFile.startsWith(ResourceUtils.CLASSPATH_URL_PREFIX)) {
Path resolved = Path.of(storeFile).toAbsolutePath();
if (Files.exists(resolved)) {
return Collections.singletonList(resolved);
}
// Include the path even if the file doesn't exist yet — the watcher uses mtime=0 / checksum="" as
// baseline, so a late-appearing file (e.g., mounted after boot) will be detected and trigger a reload.
return Collections.singletonList(Path.of(storeFile).toAbsolutePath());
}
return Collections.emptyList();
}

8
common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/config/ssl/PemSslCredentials.java

@ -33,7 +33,6 @@ import org.thingsboard.server.common.data.StringUtils;
import java.io.IOException;
import java.io.InputStream;
import java.io.InputStreamReader;
import java.nio.file.Files;
import java.nio.file.Path;
import java.security.GeneralSecurityException;
import java.security.KeyStore;
@ -152,10 +151,9 @@ public class PemSslCredentials extends AbstractSslCredentials {
private static void addIfFileSystemPath(List<Path> paths, String filePath) {
if (!StringUtils.isEmpty(filePath) && !filePath.startsWith(ResourceUtils.CLASSPATH_URL_PREFIX)) {
Path resolved = Path.of(filePath).toAbsolutePath();
if (Files.exists(resolved)) {
paths.add(resolved);
}
// Include the path even if the file doesn't exist yet — the watcher uses mtime=0 / checksum="" as
// baseline, so a late-appearing file (e.g. mounted after boot) will be detected and trigger a reload.
paths.add(Path.of(filePath).toAbsolutePath());
}
}

23
common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/config/ssl/SslCredentialsConfig.java

@ -68,20 +68,23 @@ public class SslCredentialsConfig {
}
public void onCertificateFileChanged() {
log.info("{}: Certificate file changed. Reloading SSL credentials...", name);
try {
log.info("{}: Certificate file changed. Reloading SSL credentials...", name);
this.credentials.reload(this.trustsOnly);
log.info("{}: SSL credentials reloaded successfully.", name);
for (Runnable callback : reloadCallbacks) {
try {
callback.run();
} catch (Exception e) {
log.error("{}: Error executing reload callback", name, e);
}
}
} catch (Exception e) {
log.error("{}: Failed to reload SSL credentials", name, e);
// Rethrow, so CertificateReloadManager's watcher counts this as a failure
// and applies MAX_CONSECUTIVE_FAILURES backoff instead of treating it as a successful reload.
throw new RuntimeException(name + ": Failed to reload SSL credentials", e);
}
log.info("{}: SSL credentials reloaded successfully.", name);
for (Runnable callback : reloadCallbacks) {
try {
callback.run();
} catch (Exception e) {
log.error("{}: Error executing reload callback", name, e);
}
}
}

37
common/transport/transport-api/src/main/java/org/thingsboard/server/common/transport/service/CertificateReloadManager.java

@ -32,6 +32,7 @@ import java.io.InputStream;
import java.nio.file.Files;
import java.nio.file.Path;
import java.security.MessageDigest;
import java.util.ArrayList;
import java.util.Base64;
import java.util.HashMap;
import java.util.List;
@ -48,7 +49,7 @@ public class CertificateReloadManager implements SmartInitializingSingleton, Dis
private static final int MAX_CONSECUTIVE_FAILURES = 10;
@Value("${transport.ssl.certificate.reload.enabled:false}")
@Value("${transport.ssl.certificate.reload.enabled:true}")
private boolean reloadEnabled;
@Value("${transport.ssl.certificate.reload.check_interval_seconds:60}")
@ -107,19 +108,24 @@ public class CertificateReloadManager implements SmartInitializingSingleton, Dis
continue;
}
List<Path> existingPaths = filePaths.stream()
.filter(p -> p != null && Files.exists(p))
.toList();
// Register all configured paths, including those that don't exist yet — the watcher uses
// mtime=0 / checksum="" as baseline, so files that appear later (e.g. delayed mounts) are
// picked up and trigger a reload on the next poll.
List<Path> pathsToWatch = new ArrayList<>(filePaths.size());
for (Path filePath : filePaths) {
if (filePath == null || !Files.exists(filePath)) {
log.warn("Certificate file does not exist: {} (from {})", filePath, config.getName());
if (filePath == null) {
continue;
}
pathsToWatch.add(filePath);
if (!Files.exists(filePath)) {
log.warn("Certificate file does not exist yet: {} (from {}) — will be watched and picked up when it appears",
filePath, config.getName());
}
}
if (!existingPaths.isEmpty()) {
registerWatcher(config.getName(), existingPaths, config::onCertificateFileChanged);
log.info("Registered certificate watcher: {} -> {}", config.getName(), existingPaths);
if (!pathsToWatch.isEmpty()) {
registerWatcher(config.getName(), pathsToWatch, config::onCertificateFileChanged);
log.info("Registered certificate watcher: {} -> {}", config.getName(), pathsToWatch);
}
} catch (Exception e) {
@ -190,10 +196,13 @@ public class CertificateReloadManager implements SmartInitializingSingleton, Dis
return;
}
// Compute combined checksum of all files
// Capture mtimes and checksums together before the callback runs.
// Pairing a post-callback mtime with a pre-callback checksum would let a write-during-reload be missed on the next poll.
Map<Path, Long> currentModifiedTimes = new HashMap<>();
Map<Path, String> currentChecksums = new HashMap<>();
StringBuilder combined = new StringBuilder();
for (Path path : paths) {
currentModifiedTimes.put(path, getLastModifiedTime(path));
String checksum = calculateChecksum(path);
currentChecksums.put(path, checksum);
if (!combined.isEmpty()) {
@ -216,7 +225,7 @@ public class CertificateReloadManager implements SmartInitializingSingleton, Dis
if (combinedChecksum.equals(oldCombinedChecksum)) {
// Content unchanged, just update modification times
for (Path path : paths) {
lastModifiedMap.put(path, getLastModifiedTime(path));
lastModifiedMap.put(path, currentModifiedTimes.get(path));
}
return;
}
@ -230,7 +239,7 @@ public class CertificateReloadManager implements SmartInitializingSingleton, Dis
if (consecutiveFailures >= MAX_CONSECUTIVE_FAILURES) {
// Update modification times to avoid re-checking mtime and re-computing checksums every poll cycle
for (Path path : paths) {
lastModifiedMap.put(path, getLastModifiedTime(path));
lastModifiedMap.put(path, currentModifiedTimes.get(path));
}
return;
}
@ -239,7 +248,7 @@ public class CertificateReloadManager implements SmartInitializingSingleton, Dis
log.info("Certificate change detected for: {}. Triggering reload...", name);
reloadCallback.run();
for (Path path : paths) {
lastModifiedMap.put(path, getLastModifiedTime(path));
lastModifiedMap.put(path, currentModifiedTimes.get(path));
lastChecksumMap.put(path, currentChecksums.get(path));
}
consecutiveFailures = 0;

7
common/transport/transport-api/src/test/java/org/thingsboard/server/common/transport/config/ssl/SslCredentialsConfigTest.java

@ -26,6 +26,7 @@ import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicInteger;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatThrownBy;
import static org.mockito.Mockito.doNothing;
import static org.mockito.Mockito.doThrow;
import static org.mockito.Mockito.verify;
@ -114,7 +115,7 @@ public class SslCredentialsConfigTest {
}
@Test
public void givenCredentialsReloadFails_whenCertificateChanged_thenCallbacksShouldNotBeCalled() throws Exception {
public void givenCredentialsReloadFails_whenCertificateChanged_thenShouldRethrowAndNotCallCallbacks() throws Exception {
AtomicInteger callbackCount = new AtomicInteger(0);
config.registerReloadCallback(callbackCount::incrementAndGet);
@ -122,7 +123,9 @@ public class SslCredentialsConfigTest {
doThrow(new RuntimeException("Simulated reload failure")).when(mockCredentials).reload(false);
config.onCertificateFileChanged();
assertThatThrownBy(() -> config.onCertificateFileChanged())
.isInstanceOf(RuntimeException.class)
.hasMessageContaining("Failed to reload SSL credentials");
assertThat(callbackCount.get()).isEqualTo(0);
}

Loading…
Cancel
Save