diff --git a/rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroup.java b/rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroup.java index 0a099df50e..9d15cb9793 100644 --- a/rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroup.java +++ b/rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroup.java @@ -26,14 +26,18 @@ import org.thingsboard.common.util.SsrfProtectionValidator; import java.net.InetAddress; import java.net.InetSocketAddress; import java.net.SocketAddress; +import java.util.ArrayList; +import java.util.HashSet; import java.util.List; -import java.util.stream.Collectors; +import java.util.Set; /** * Custom Netty {@link AddressResolverGroup} that validates every resolved IP address * against the SSRF block-list at connection time. This eliminates the DNS rebinding * TOCTOU gap where a hostname resolves to a safe IP during validation but to a * private/metadata IP when the actual connection is made. + *

+ * Only wired into {@link TbHttpClient} when SSRF protection is enabled. */ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup { @@ -76,16 +80,22 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup resolve(SocketAddress address, Promise promise) { delegate.resolve(address).addListener((Future future) -> { - if (!future.isSuccess()) { - promise.tryFailure(future.cause()); - return; - } - InetSocketAddress resolved = future.getNow(); - if (SsrfProtectionValidator.isEnabled() && isBlocked(resolved) && !isOriginalHostAllowed(address)) { - promise.tryFailure(new RuntimeException( - "URI is invalid: host '" + resolved.getAddress().getHostAddress() + "' is not allowed")); - } else { - promise.trySuccess(resolved); + try { + if (!future.isSuccess()) { + promise.tryFailure(future.cause()); + return; + } + InetSocketAddress resolved = future.getNow(); + if (isOriginalHostAllowed(address)) { + promise.trySuccess(resolved); + } else if (isBlocked(resolved)) { + promise.tryFailure(new RuntimeException( + "URI is invalid: host '" + getHostString(address) + "' is not allowed")); + } else { + promise.trySuccess(resolved); + } + } catch (Exception e) { + promise.tryFailure(e); } }); return promise; @@ -99,24 +109,41 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup> resolveAll(SocketAddress address, Promise> promise) { delegate.resolveAll(address).addListener((Future> future) -> { - if (!future.isSuccess()) { - promise.tryFailure(future.cause()); - return; - } - List resolved = future.getNow(); - if (!SsrfProtectionValidator.isEnabled() || isOriginalHostAllowed(address)) { - promise.trySuccess(resolved); - return; - } - List safe = resolved.stream() - .filter(addr -> !isBlocked(addr)) - .collect(Collectors.toList()); - if (safe.isEmpty()) { - String host = address instanceof InetSocketAddress isa ? isa.getHostString() : address.toString(); - promise.tryFailure(new RuntimeException( - "URI is invalid: host '" + host + "' is not allowed")); - } else { - promise.trySuccess(safe); + try { + if (!future.isSuccess()) { + promise.tryFailure(future.cause()); + return; + } + List resolved = future.getNow(); + if (isOriginalHostAllowed(address)) { + promise.trySuccess(resolved); + return; + } + Set blocked = null; + for (InetSocketAddress addr : resolved) { + if (isBlocked(addr)) { + if (blocked == null) { + blocked = new HashSet<>(2); + } + blocked.add(addr); + } + } + if (blocked == null) { + promise.trySuccess(resolved); + } else if (blocked.size() == resolved.size()) { + promise.tryFailure(new RuntimeException( + "URI is invalid: host '" + getHostString(address) + "' is not allowed")); + } else { + List safe = new ArrayList<>(resolved.size() - blocked.size()); + for (InetSocketAddress addr : resolved) { + if (!blocked.contains(addr)) { + safe.add(addr); + } + } + promise.trySuccess(safe); + } + } catch (Exception e) { + promise.tryFailure(e); } }); return promise; @@ -139,5 +166,9 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup resolver = SsrfSafeAddressResolverGroup.INSTANCE.getResolver(executor); Promise> promise = executor.newPromise(); - executor.submit(() -> resolver.resolveAll(InetSocketAddress.createUnresolved("127.0.0.1", 80), promise)); + executor.submit(() -> resolver.resolveAll(InetSocketAddress.createUnresolved("8.8.8.8", 80), promise)); List results = promise.get(10, TimeUnit.SECONDS); assertThat(results).isNotEmpty(); + assertThat(results.get(0).getAddress().getHostAddress()).isEqualTo("8.8.8.8"); + } + + @Test + void resolveAllPrivateIpFailsWhenSsrfEnabled() { + assertThatThrownBy(() -> { + EventExecutor executor = eventLoopGroup.next(); + AddressResolver resolver = SsrfSafeAddressResolverGroup.INSTANCE.getResolver(executor); + Promise> promise = executor.newPromise(); + executor.submit(() -> resolver.resolveAll(InetSocketAddress.createUnresolved("127.0.0.1", 80), promise)); + promise.get(10, TimeUnit.SECONDS); + }).isInstanceOf(ExecutionException.class) + .hasRootCauseInstanceOf(RuntimeException.class) + .rootCause().hasMessageContaining("is not allowed"); } }