diff --git a/plugins/integrations/kubernetes-service/src/main/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorker.java b/plugins/integrations/kubernetes-service/src/main/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorker.java index 55924cb32e8a..f2fbdbc6f53f 100644 --- a/plugins/integrations/kubernetes-service/src/main/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorker.java +++ b/plugins/integrations/kubernetes-service/src/main/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorker.java @@ -544,16 +544,23 @@ protected FirewallRule removeApiFirewallRule(final IpAddress publicIp) { } protected FirewallRule removeSshFirewallRule(final IpAddress publicIp, final long networkId) { - FirewallRule rule = null; + FirewallRuleVO rule = null; List firewallRules = firewallRulesDao.listByIpPurposeProtocolAndNotRevoked(publicIp.getId(), FirewallRule.Purpose.Firewall, NetUtils.TCP_PROTO); for (FirewallRuleVO firewallRule : firewallRules) { - PortForwardingRuleVO pfRule = portForwardingRulesDao.findByNetworkAndPorts(networkId, firewallRule.getSourcePortStart(), firewallRule.getSourcePortEnd()); - if (Objects.equals(firewallRule.getSourcePortStart(), CLUSTER_NODES_DEFAULT_START_SSH_PORT) || (Objects.nonNull(pfRule) && pfRule.getDestinationPortStart() == DEFAULT_SSH_PORT) ) { + if (Objects.equals(firewallRule.getSourcePortStart(), CLUSTER_NODES_DEFAULT_START_SSH_PORT)) { rule = firewallRule; - firewallService.revokeIngressFwRule(firewallRule.getId(), true); - logger.debug("The SSH firewall rule {} with the id {} was revoked", firewallRule.getName(), firewallRule.getId()); break; } + if (rule == null) { + PortForwardingRuleVO pfRule = portForwardingRulesDao.findByNetworkAndPorts(networkId, firewallRule.getSourcePortStart(), firewallRule.getSourcePortEnd()); + if (Objects.nonNull(pfRule) && pfRule.getDestinationPortStart() == DEFAULT_SSH_PORT) { + rule = firewallRule; + } + } + } + if (rule != null) { + firewallService.revokeIngressFwRule(rule.getId(), true); + logger.debug("The SSH firewall rule {} with the id {} was revoked", rule.getName(), rule.getId()); } return rule; } diff --git a/plugins/integrations/kubernetes-service/src/test/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorkerTest.java b/plugins/integrations/kubernetes-service/src/test/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorkerTest.java index c220a3468afb..597fa6974050 100644 --- a/plugins/integrations/kubernetes-service/src/test/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorkerTest.java +++ b/plugins/integrations/kubernetes-service/src/test/java/com/cloud/kubernetes/cluster/actionworkers/KubernetesClusterResourceModifierActionWorkerTest.java @@ -23,6 +23,15 @@ import com.cloud.kubernetes.cluster.dao.KubernetesClusterDetailsDao; import com.cloud.kubernetes.cluster.dao.KubernetesClusterVmMapDao; import com.cloud.kubernetes.version.dao.KubernetesSupportedVersionDao; +import com.cloud.network.IpAddress; +import com.cloud.network.dao.FirewallRulesDao; +import com.cloud.network.firewall.FirewallService; +import com.cloud.network.rules.FirewallRule; +import com.cloud.network.rules.FirewallRuleVO; +import com.cloud.network.rules.PortForwardingRuleVO; +import com.cloud.network.rules.dao.PortForwardingRulesDao; +import com.cloud.utils.net.NetUtils; +import java.util.Arrays; import org.junit.Assert; import org.junit.Before; import org.junit.Test; @@ -51,6 +60,18 @@ public class KubernetesClusterResourceModifierActionWorkerTest { @Mock private KubernetesCluster kubernetesClusterMock; + @Mock + private IpAddress publicIpMock; + + @Mock + private FirewallRulesDao firewallRulesDaoMock; + + @Mock + private PortForwardingRulesDao portForwardingRulesDaoMock; + + @Mock + private FirewallService firewallServiceMock; + private KubernetesClusterResourceModifierActionWorker kubernetesClusterResourceModifierActionWorker; @Before @@ -135,4 +156,58 @@ public void getKubernetesClusterNodeNamePrefixTestNormalizedPrefixShouldNotStart Mockito.when(kubernetesClusterMock.getName()).thenReturn(originalPrefix); Assert.assertEquals(expectedPrefix, kubernetesClusterResourceModifierActionWorker.getKubernetesClusterNodeNamePrefix()); } + + private static final long NETWORK_ID = 10L; + + private void mockSshFirewallRules(FirewallRuleVO... rules) { + kubernetesClusterResourceModifierActionWorker.firewallRulesDao = firewallRulesDaoMock; + kubernetesClusterResourceModifierActionWorker.portForwardingRulesDao = portForwardingRulesDaoMock; + kubernetesClusterResourceModifierActionWorker.firewallService = firewallServiceMock; + Mockito.when(publicIpMock.getId()).thenReturn(1L); + Mockito.when(firewallRulesDaoMock.listByIpPurposeProtocolAndNotRevoked(1L, FirewallRule.Purpose.Firewall, NetUtils.TCP_PROTO)).thenReturn(Arrays.asList(rules)); + } + + private FirewallRuleVO mockSshForwardedFirewallRule(int port) { + FirewallRuleVO rule = Mockito.mock(FirewallRuleVO.class); + Mockito.when(rule.getSourcePortStart()).thenReturn(port); + Mockito.when(rule.getSourcePortEnd()).thenReturn(port); + PortForwardingRuleVO pfRule = Mockito.mock(PortForwardingRuleVO.class); + Mockito.when(pfRule.getDestinationPortStart()).thenReturn(KubernetesClusterActionWorker.DEFAULT_SSH_PORT); + Mockito.when(portForwardingRulesDaoMock.findByNetworkAndPorts(NETWORK_ID, port, port)).thenReturn(pfRule); + return rule; + } + + @Test + public void removeSshFirewallRuleTestPrefersNodesRuleOverEtcdRuleListedFirst() { + FirewallRuleVO etcdRule = mockSshForwardedFirewallRule(50000); + FirewallRuleVO nodesRule = Mockito.mock(FirewallRuleVO.class); + Mockito.when(nodesRule.getSourcePortStart()).thenReturn(KubernetesClusterActionWorker.CLUSTER_NODES_DEFAULT_START_SSH_PORT); + Mockito.when(nodesRule.getId()).thenReturn(11L); + mockSshFirewallRules(etcdRule, nodesRule); + + Assert.assertSame(nodesRule, kubernetesClusterResourceModifierActionWorker.removeSshFirewallRule(publicIpMock, NETWORK_ID)); + Mockito.verify(firewallServiceMock).revokeIngressFwRule(11L, true); + Mockito.verifyNoMoreInteractions(firewallServiceMock); + } + + @Test + public void removeSshFirewallRuleTestFallsBackToSshForwardedRule() { + FirewallRuleVO externalNodeRule = mockSshForwardedFirewallRule(2225); + Mockito.when(externalNodeRule.getId()).thenReturn(12L); + mockSshFirewallRules(externalNodeRule); + + Assert.assertSame(externalNodeRule, kubernetesClusterResourceModifierActionWorker.removeSshFirewallRule(publicIpMock, NETWORK_ID)); + Mockito.verify(firewallServiceMock).revokeIngressFwRule(12L, true); + } + + @Test + public void removeSshFirewallRuleTestReturnsNullWhenNoSshRule() { + FirewallRuleVO apiRule = Mockito.mock(FirewallRuleVO.class); + Mockito.when(apiRule.getSourcePortStart()).thenReturn(KubernetesClusterActionWorker.CLUSTER_API_PORT); + Mockito.when(apiRule.getSourcePortEnd()).thenReturn(KubernetesClusterActionWorker.CLUSTER_API_PORT); + mockSshFirewallRules(apiRule); + + Assert.assertNull(kubernetesClusterResourceModifierActionWorker.removeSshFirewallRule(publicIpMock, NETWORK_ID)); + Mockito.verifyNoMoreInteractions(firewallServiceMock); + } }