Browse Source

Optimize SsrfSafeAddressResolverGroup, remove dead isEnabled checks

Replace stream().filter().collect() in resolveAll with single-pass loop
to avoid allocations on the common path (nothing blocked). Remove
redundant isEnabled() checks inside the resolver since it is only wired
when SSRF protection is enabled. Add resolveAll test coverage.
pull/15253/head
Viacheslav Klimov 6 months ago
parent
commit
d83a28beaa
Failed to extract signature
  1. 89
      rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroup.java
  2. 20
      rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroupTest.java

89
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.
* <p>
* Only wired into {@link TbHttpClient} when SSRF protection is enabled.
*/
public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup<InetSocketAddress> {
@ -76,16 +80,22 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup<Ine
@Override
public Future<InetSocketAddress> resolve(SocketAddress address, Promise<InetSocketAddress> promise) {
delegate.resolve(address).addListener((Future<InetSocketAddress> 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<Ine
@Override
public Future<List<InetSocketAddress>> resolveAll(SocketAddress address, Promise<List<InetSocketAddress>> promise) {
delegate.resolveAll(address).addListener((Future<List<InetSocketAddress>> future) -> {
if (!future.isSuccess()) {
promise.tryFailure(future.cause());
return;
}
List<InetSocketAddress> resolved = future.getNow();
if (!SsrfProtectionValidator.isEnabled() || isOriginalHostAllowed(address)) {
promise.trySuccess(resolved);
return;
}
List<InetSocketAddress> 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<InetSocketAddress> resolved = future.getNow();
if (isOriginalHostAllowed(address)) {
promise.trySuccess(resolved);
return;
}
Set<InetSocketAddress> 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<InetSocketAddress> 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<Ine
}
return false;
}
private static String getHostString(SocketAddress address) {
return address instanceof InetSocketAddress isa ? isa.getHostString() : address.toString();
}
}
}

20
rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroupTest.java

@ -133,17 +133,29 @@ class SsrfSafeAddressResolverGroupTest {
}
@Test
void resolveAllReturnsAllWhenSsrfDisabled() throws Exception {
SsrfProtectionValidator.setEnabled(false);
void resolveAllPublicIpSucceeds() throws Exception {
EventExecutor executor = eventLoopGroup.next();
AddressResolver<InetSocketAddress> resolver = SsrfSafeAddressResolverGroup.INSTANCE.getResolver(executor);
Promise<List<InetSocketAddress>> 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<InetSocketAddress> 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<InetSocketAddress> resolver = SsrfSafeAddressResolverGroup.INSTANCE.getResolver(executor);
Promise<List<InetSocketAddress>> 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");
}
}

Loading…
Cancel
Save