diff --git a/application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/CustomOAuth2ClientMapper.java b/application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/CustomOAuth2ClientMapper.java index 5cd6ca09b2..8477c69a99 100644 --- a/application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/CustomOAuth2ClientMapper.java +++ b/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) { diff --git a/application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/GithubOAuth2ClientMapper.java b/application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/GithubOAuth2ClientMapper.java index 2861036097..c7aa893d59 100644 --- a/application/src/main/java/org/thingsboard/server/service/security/auth/oauth2/GithubOAuth2ClientMapper.java +++ b/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(); diff --git a/common/util/src/main/java/org/thingsboard/common/util/SsrfProtectionValidator.java b/common/util/src/main/java/org/thingsboard/common/util/SsrfProtectionValidator.java index eb89b164a2..9132d117ae 100644 --- a/common/util/src/main/java/org/thingsboard/common/util/SsrfProtectionValidator.java +++ b/common/util/src/main/java/org/thingsboard/common/util/SsrfProtectionValidator.java @@ -167,9 +167,28 @@ public class SsrfProtectionValidator { } public static void setAdditionalBlockedHosts(List 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 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 entries) { if (entries == null || entries.isEmpty()) { - additionalBlocked = AdditionalBlockedHosts.EMPTY; - return; + return ParsedHostEntries.EMPTY; } List cidrRanges = new ArrayList<>(); Set 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 entries) { - if (entries == null || entries.isEmpty()) { - allowedHosts = AllowedHosts.EMPTY; - return; - } - List cidrRanges = new ArrayList<>(); - Set 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 cidrRanges, Set hostnames) { + static final ParsedHostEntries EMPTY = new ParsedHostEntries(Collections.emptyList(), Collections.emptySet()); } record AdditionalBlockedHosts(List cidrRanges, Set hostnames) { diff --git a/dao/src/main/java/org/thingsboard/server/dao/service/validator/Oauth2ClientDataValidator.java b/dao/src/main/java/org/thingsboard/server/dao/service/validator/Oauth2ClientDataValidator.java index 07fbc06114..d5b965face 100644 --- a/dao/src/main/java/org/thingsboard/server/dao/service/validator/Oauth2ClientDataValidator.java +++ b/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 { @@ -64,6 +67,11 @@ public class Oauth2ClientDataValidator extends DataValidator { 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()); + } } } } 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 f3cd0c83bb..e08dc0ac06 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 @@ -81,7 +81,7 @@ public final class SsrfSafeAddressResolverGroup extends AddressResolverGroup 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 resolver = SsrfSafeAddressResolverGroup.INSTANCE.getResolver(executor); Promise 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