From b26d56e25a32134505d739db0f905602b449b6e0 Mon Sep 17 00:00:00 2001 From: Alice Date: Mon, 3 Aug 2026 03:03:56 +0700 Subject: [PATCH] fix: scope GCP firewall rules per node --- deploy_gcp/seismic_deploy/gcp/compute.py | 17 +++- deploy_gcp/seismic_deploy/gcp/test_compute.py | 79 +++++++++++++++++++ 2 files changed, 93 insertions(+), 3 deletions(-) create mode 100644 deploy_gcp/seismic_deploy/gcp/test_compute.py diff --git a/deploy_gcp/seismic_deploy/gcp/compute.py b/deploy_gcp/seismic_deploy/gcp/compute.py index aba7192f..37565c63 100644 --- a/deploy_gcp/seismic_deploy/gcp/compute.py +++ b/deploy_gcp/seismic_deploy/gcp/compute.py @@ -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) @@ -296,7 +302,7 @@ 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) @@ -304,3 +310,8 @@ def delete_firewall_rules(project: str, node_name: str) -> None: 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}" diff --git a/deploy_gcp/seismic_deploy/gcp/test_compute.py b/deploy_gcp/seismic_deploy/gcp/test_compute.py new file mode 100644 index 00000000..6c341003 --- /dev/null +++ b/deploy_gcp/seismic_deploy/gcp/test_compute.py @@ -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()