Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 14 additions & 3 deletions deploy_gcp/seismic_deploy/gcp/compute.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,11 +87,17 @@ def _create():


def ensure_firewall_rules(project: str, node_name: str) -> None:
"""Create firewall rules for the node. Idempotent."""
"""Create firewall rules scoped to this node. Idempotent.

Firewall target tags are node-specific, so the rule name must be
node-specific as well. Reusing a global rule name would make subsequent
nodes skip creation even though the existing rule targets only the first
node.
"""
client = compute_v1.FirewallsClient()

for base_name, proto, port in FIREWALL_RULES:
rule_name = f"{base_name}"
rule_name = _firewall_rule_name(base_name, node_name)

try:
client.get(project=project, firewall=rule_name)
Expand Down Expand Up @@ -296,11 +302,16 @@ def delete_firewall_rules(project: str, node_name: str) -> None:
"""Delete all firewall rules for a node."""
client = compute_v1.FirewallsClient()
for base_name, _, _ in FIREWALL_RULES:
rule_name = f"{base_name}-{node_name}"
rule_name = _firewall_rule_name(base_name, node_name)
try:
click.echo(f" Deleting firewall rule: {rule_name}...")
op = client.delete(project=project, firewall=rule_name)
op.result()
click.echo(f" Firewall rule deleted: {rule_name}")
except NotFound:
pass


def _firewall_rule_name(base_name: str, node_name: str) -> str:
"""Return the stable per-node GCP firewall rule name."""
return f"{base_name}-{node_name}"
79 changes: 79 additions & 0 deletions deploy_gcp/seismic_deploy/gcp/test_compute.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
"""Regression tests for GCP firewall rule lifecycle."""

from __future__ import annotations

import unittest
from typing import Any
from unittest.mock import patch

from google.api_core.exceptions import NotFound

from deploy_gcp.seismic_deploy.gcp.compute import (
FIREWALL_RULES,
delete_firewall_rules,
ensure_firewall_rules,
)


class _Operation:
def result(self) -> None:
return None


class _FirewallClient:
def __init__(self) -> None:
self.rules: dict[str, object] = {}
self.inserted: list[Any] = []
self.deleted: list[str] = []

def get(self, *, project: str, firewall: str) -> object:
del project
try:
return self.rules[firewall]
except KeyError:
raise NotFound(firewall) from None

def insert(self, *, project: str, firewall_resource: Any) -> _Operation:
del project
name = firewall_resource.name
self.rules[name] = firewall_resource
self.inserted.append(firewall_resource)
return _Operation()

def delete(self, *, project: str, firewall: str) -> _Operation:
del project
self.deleted.append(firewall)
self.rules.pop(firewall, None)
return _Operation()


class FirewallRuleLifecycleTests(unittest.TestCase):
@patch("deploy_gcp.seismic_deploy.gcp.compute.compute_v1.FirewallsClient")
def test_rules_are_unique_per_node(self, client_cls) -> None:
client = _FirewallClient()
client_cls.return_value = client

ensure_firewall_rules("project", "node-a")
ensure_firewall_rules("project", "node-b")

self.assertEqual(len(client.inserted), len(FIREWALL_RULES) * 2)
node_a = client.inserted[: len(FIREWALL_RULES)]
node_b = client.inserted[len(FIREWALL_RULES) :]
self.assertTrue(all(rule.name.endswith("-node-a") for rule in node_a))
self.assertTrue(all(rule.name.endswith("-node-b") for rule in node_b))
self.assertTrue(all(rule.target_tags == ["node-a"] for rule in node_a))
self.assertTrue(all(rule.target_tags == ["node-b"] for rule in node_b))

@patch("deploy_gcp.seismic_deploy.gcp.compute.compute_v1.FirewallsClient")
def test_destroy_deletes_the_same_per_node_names(self, client_cls) -> None:
client = _FirewallClient()
client_cls.return_value = client

delete_firewall_rules("project", "node-a")

expected = [f"{base}-node-a" for base, _, _ in FIREWALL_RULES]
self.assertEqual(client.deleted, expected)


if __name__ == "__main__":
unittest.main()