Browse Source

Address PR review comments

- Extract shared parseHostEntries() to deduplicate setAllowedHosts/setAdditionalBlockedHosts
- Add isHostnameAllowed() and propagate hostname allow-list check in resolver
- Move OAuth2 custom mapper URL SSRF validation to save-time (Oauth2ClientDataValidator)
- Remove runtime SSRF checks from CustomOAuth2ClientMapper and GithubOAuth2ClientMapper
  (custom URL now validated at save; GitHub emailUrl is server config, not user input)
- Replace example.com with 8.8.8.8 in resolver test to avoid DNS dependency
pull/15253/head
Viacheslav Klimov 7 months ago
parent
commit
2f39347dd2
Failed to extract signature
  1. 9
      application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/CustomOAuth2ClientMapper.java
  2. 8
      application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/GithubOAuth2ClientMapper.java
  3. 54
      common/util/src/main/java/org/thingsboard/common/util/SsrfProtectionValidator.java
  4. 8
      dao/src/main/java/org/thingsboard/server/dao/service/validator/Oauth2ClientDataValidator.java
  5. 12
      rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroup.java
  6. 5
      rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroupTest.java

9
application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/CustomOAuth2ClientMapper.java

@ -23,14 +23,11 @@ import org.springframework.security.oauth2.client.authentication.OAuth2Authentic
import org.springframework.stereotype.Service;
import org.springframework.web.client.RestTemplate;
import org.thingsboard.common.util.JacksonUtil;
import org.thingsboard.common.util.SsrfProtectionValidator;
import org.thingsboard.server.common.data.StringUtils;
import org.thingsboard.server.common.data.oauth2.OAuth2CustomMapperConfig;
import org.thingsboard.server.common.data.oauth2.OAuth2MapperConfig;
import org.thingsboard.server.common.data.oauth2.OAuth2Client;
import org.thingsboard.server.dao.oauth2.OAuth2User;
import java.net.URI;
import org.thingsboard.server.queue.util.TbCoreComponent;
import org.thingsboard.server.service.security.model.SecurityUser;
@ -66,12 +63,6 @@ public class CustomOAuth2ClientMapper extends AbstractOAuth2ClientMapper impleme
log.error("Can't convert principal to JSON string", e);
throw new RuntimeException("Can't convert principal to JSON string", e);
}
try {
SsrfProtectionValidator.validateUri(new URI(custom.getUrl()));
} catch (Exception e) {
log.error("SSRF validation failed for custom mapper URL '{}'", custom.getUrl(), e);
throw new RuntimeException("Unable to login. Please contact your Administrator!");
}
try {
return restTemplate.postForEntity(custom.getUrl(), request, OAuth2User.class).getBody();
} catch (Exception e) {

8
application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/GithubOAuth2ClientMapper.java

@ -24,7 +24,6 @@ import org.springframework.boot.web.client.RestTemplateBuilder;
import org.springframework.security.oauth2.client.authentication.OAuth2AuthenticationToken;
import org.springframework.stereotype.Service;
import org.springframework.web.client.RestTemplate;
import org.thingsboard.common.util.SsrfProtectionValidator;
import org.thingsboard.server.common.data.oauth2.OAuth2MapperConfig;
import org.thingsboard.server.common.data.oauth2.OAuth2Client;
import org.thingsboard.server.dao.oauth2.OAuth2Configuration;
@ -32,7 +31,6 @@ import org.thingsboard.server.dao.oauth2.OAuth2User;
import org.thingsboard.server.queue.util.TbCoreComponent;
import org.thingsboard.server.service.security.model.SecurityUser;
import java.net.URI;
import java.util.ArrayList;
import java.util.Map;
import java.util.Optional;
@ -64,12 +62,6 @@ public class GithubOAuth2ClientMapper extends AbstractOAuth2ClientMapper impleme
restTemplateBuilder = restTemplateBuilder.defaultHeader(AUTHORIZATION, "token " + oauth2Token);
RestTemplate restTemplate = restTemplateBuilder.build();
try {
SsrfProtectionValidator.validateUri(new URI(emailUrl));
} catch (Exception e) {
log.error("SSRF validation failed for GitHub email URL '{}'", emailUrl, e);
throw new RuntimeException("Unable to login. Please contact your Administrator!");
}
GithubEmailsResponse githubEmailsResponse;
try {
githubEmailsResponse = restTemplate.getForEntity(emailUrl, GithubEmailsResponse.class).getBody();

54
common/util/src/main/java/org/thingsboard/common/util/SsrfProtectionValidator.java

@ -167,9 +167,28 @@ public class SsrfProtectionValidator {
}
public static void setAdditionalBlockedHosts(List<String> entries) {
ParsedHostEntries parsed = parseHostEntries(entries);
additionalBlocked = new AdditionalBlockedHosts(parsed.cidrRanges, parsed.hostnames);
if (!parsed.cidrRanges.isEmpty() || !parsed.hostnames.isEmpty()) {
log.info("SSRF additional blocked hosts configured: {} CIDR range(s), {} hostname(s)", parsed.cidrRanges.size(), parsed.hostnames.size());
}
}
public static void setAllowedHosts(List<String> entries) {
ParsedHostEntries parsed = parseHostEntries(entries);
allowedHosts = new AllowedHosts(parsed.cidrRanges, parsed.hostnames);
if (!parsed.cidrRanges.isEmpty() || !parsed.hostnames.isEmpty()) {
log.info("SSRF allowed hosts configured: {} CIDR range(s), {} hostname(s)", parsed.cidrRanges.size(), parsed.hostnames.size());
}
}
public static boolean isHostnameAllowed(String hostname) {
return allowedHosts.hostnames.contains(hostname.toLowerCase());
}
private static ParsedHostEntries parseHostEntries(List<String> entries) {
if (entries == null || entries.isEmpty()) {
additionalBlocked = AdditionalBlockedHosts.EMPTY;
return;
return ParsedHostEntries.EMPTY;
}
List<CidrRange> cidrRanges = new ArrayList<>();
Set<String> hostnames = new HashSet<>();
@ -188,10 +207,9 @@ public class SsrfProtectionValidator {
hostnames.add(trimmed.toLowerCase());
}
}
additionalBlocked = new AdditionalBlockedHosts(
return new ParsedHostEntries(
Collections.unmodifiableList(cidrRanges),
Collections.unmodifiableSet(hostnames));
log.info("SSRF additional blocked hosts configured: {} CIDR range(s), {} hostname(s)", cidrRanges.size(), hostnames.size());
}
private static boolean isIpLiteral(String entry) {
@ -199,32 +217,8 @@ public class SsrfProtectionValidator {
return !entry.isEmpty() && (Character.isDigit(entry.charAt(0)) || entry.contains(":"));
}
public static void setAllowedHosts(List<String> entries) {
if (entries == null || entries.isEmpty()) {
allowedHosts = AllowedHosts.EMPTY;
return;
}
List<CidrRange> cidrRanges = new ArrayList<>();
Set<String> hostnames = new HashSet<>();
for (String entry : entries) {
String trimmed = entry.trim();
if (trimmed.isEmpty()) {
continue;
}
if (trimmed.contains("/") || isIpLiteral(trimmed)) {
try {
cidrRanges.add(CidrRange.parse(trimmed));
} catch (Exception e) {
log.warn("Failed to parse allowed CIDR/IP entry '{}': {}", trimmed, e.getMessage());
}
} else {
hostnames.add(trimmed.toLowerCase());
}
}
allowedHosts = new AllowedHosts(
Collections.unmodifiableList(cidrRanges),
Collections.unmodifiableSet(hostnames));
log.info("SSRF allowed hosts configured: {} CIDR range(s), {} hostname(s)", cidrRanges.size(), hostnames.size());
private record ParsedHostEntries(List<CidrRange> cidrRanges, Set<String> hostnames) {
static final ParsedHostEntries EMPTY = new ParsedHostEntries(Collections.emptyList(), Collections.emptySet());
}
record AdditionalBlockedHosts(List<CidrRange> cidrRanges, Set<String> hostnames) {

8
dao/src/main/java/org/thingsboard/server/dao/service/validator/Oauth2ClientDataValidator.java

@ -17,6 +17,7 @@ package org.thingsboard.server.dao.service.validator;
import lombok.AllArgsConstructor;
import org.springframework.stereotype.Component;
import org.thingsboard.common.util.SsrfProtectionValidator;
import org.thingsboard.server.common.data.StringUtils;
import org.thingsboard.server.common.data.id.TenantId;
import org.thingsboard.server.common.data.oauth2.MapperType;
@ -28,6 +29,8 @@ import org.thingsboard.server.common.data.oauth2.TenantNameStrategyType;
import org.thingsboard.server.dao.exception.DataValidationException;
import org.thingsboard.server.dao.service.DataValidator;
import java.net.URI;
@Component
@AllArgsConstructor
public class Oauth2ClientDataValidator extends DataValidator<OAuth2Client> {
@ -64,6 +67,11 @@ public class Oauth2ClientDataValidator extends DataValidator<OAuth2Client> {
if (StringUtils.isEmpty(customConfig.getUrl())) {
throw new DataValidationException("Custom mapper URL should be specified!");
}
try {
SsrfProtectionValidator.validateUri(new URI(customConfig.getUrl()));
} catch (Exception e) {
throw new DataValidationException("Custom mapper URL is not allowed: " + e.getMessage());
}
}
}
}

12
rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/rest/SsrfSafeAddressResolverGroup.java

@ -81,7 +81,7 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup<Ine
return;
}
InetSocketAddress resolved = future.getNow();
if (SsrfProtectionValidator.isEnabled() && isBlocked(resolved)) {
if (SsrfProtectionValidator.isEnabled() && isBlocked(resolved) && !isOriginalHostAllowed(address)) {
promise.tryFailure(new RuntimeException(
"SSRF protection: resolved address " + resolved.getAddress().getHostAddress() + " is blocked"));
} else {
@ -104,7 +104,7 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup<Ine
return;
}
List<InetSocketAddress> resolved = future.getNow();
if (!SsrfProtectionValidator.isEnabled()) {
if (!SsrfProtectionValidator.isEnabled() || isOriginalHostAllowed(address)) {
promise.trySuccess(resolved);
return;
}
@ -131,5 +131,13 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup<Ine
InetAddress addr = socketAddress.getAddress();
return addr != null && SsrfProtectionValidator.isBlockedAddress(addr);
}
private static boolean isOriginalHostAllowed(SocketAddress address) {
if (address instanceof InetSocketAddress isa) {
String host = isa.getHostString();
return host != null && SsrfProtectionValidator.isHostnameAllowed(host);
}
return false;
}
}
}

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

@ -80,12 +80,11 @@ class SsrfSafeAddressResolverGroupTest {
AddressResolver<InetSocketAddress> resolver = SsrfSafeAddressResolverGroup.INSTANCE.getResolver(executor);
Promise<InetSocketAddress> promise = executor.newPromise();
executor.submit(() -> resolver.resolve(InetSocketAddress.createUnresolved("example.com", 80), promise));
executor.submit(() -> resolver.resolve(InetSocketAddress.createUnresolved("8.8.8.8", 80), promise));
InetSocketAddress result = promise.get(10, TimeUnit.SECONDS);
assertThat(result.getAddress()).isNotNull();
assertThat(result.getAddress().isLoopbackAddress()).isFalse();
assertThat(result.getAddress().isSiteLocalAddress()).isFalse();
assertThat(result.getAddress().getHostAddress()).isEqualTo("8.8.8.8");
}
@Test

Loading…
Cancel
Save