mirror of
https://github.com/trailofbits/dropkit
synced 2026-06-21 14:11:54 +00:00
Abort hibernate when prepare_for_hibernate fails to secure SSH access
Previously, if the temporary UFW rule or SSH config update failed during hibernate preparation, the function returned True and hibernate proceeded to snapshot — producing a snapshot with no public SSH access. On wake, such droplets were completely unreachable (Tailscale needs re-auth but SSH is blocked on eth0). Now prepare_for_hibernate returns None on failure, and the caller aborts the hibernate. Also adds a warning to the --continue path since it cannot verify the temp rule on a powered-off droplet. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
+50
-28
@@ -1006,10 +1006,13 @@ def add_temporary_ssh_rule(ssh_hostname: str, verbose: bool = False) -> bool:
|
||||
if result.stderr:
|
||||
console.print(f"[dim]stderr: {result.stderr}[/dim]")
|
||||
|
||||
return result.returncode == 0
|
||||
if result.returncode != 0:
|
||||
stderr_msg = result.stderr.strip()[:200] if result.stderr else "no error output"
|
||||
console.print(f"[dim]SSH rule command failed (exit {result.returncode}): {stderr_msg}[/dim]")
|
||||
return False
|
||||
return True
|
||||
except (subprocess.TimeoutExpired, subprocess.SubprocessError) as e:
|
||||
if verbose:
|
||||
console.print(f"[dim]SSH command failed: {e}[/dim]")
|
||||
console.print(f"[dim]SSH command failed: {e}[/dim]")
|
||||
return False
|
||||
|
||||
|
||||
@@ -1019,7 +1022,7 @@ def prepare_for_hibernate(
|
||||
droplet: dict,
|
||||
droplet_name: str,
|
||||
verbose: bool = False,
|
||||
) -> bool:
|
||||
) -> bool | None:
|
||||
"""
|
||||
Prepare droplet for hibernation, handling Tailscale lockdown if present.
|
||||
|
||||
@@ -1035,7 +1038,9 @@ def prepare_for_hibernate(
|
||||
verbose: Show debug output
|
||||
|
||||
Returns:
|
||||
True if droplet was under Tailscale lockdown, False otherwise
|
||||
True if droplet was under Tailscale lockdown and prepared successfully,
|
||||
False if no Tailscale lockdown detected,
|
||||
None if preparation failed (hibernate should abort).
|
||||
"""
|
||||
if not is_droplet_tailscale_locked(config, droplet_name):
|
||||
return False
|
||||
@@ -1047,13 +1052,11 @@ def prepare_for_hibernate(
|
||||
# Add temporary SSH rule via Tailscale connection (must succeed before logout)
|
||||
console.print("[dim]Adding temporary SSH rule for eth0...[/dim]")
|
||||
if not add_temporary_ssh_rule(ssh_hostname, verbose):
|
||||
# Temp rule failed - abort logout to keep Tailscale connectivity
|
||||
console.print(
|
||||
"[yellow]⚠[/yellow] Could not add temporary SSH rule - "
|
||||
"skipping Tailscale logout to maintain connectivity"
|
||||
"[red]Error:[/red] Could not add temporary SSH rule. "
|
||||
"Aborting hibernate to prevent unreachable droplet on wake."
|
||||
)
|
||||
# Continue without logout - snapshot will preserve Tailscale state
|
||||
return True
|
||||
return None
|
||||
|
||||
console.print("[green]✓[/green] Temporary SSH rule added")
|
||||
|
||||
@@ -1069,21 +1072,30 @@ def prepare_for_hibernate(
|
||||
public_ip = network.get("ip_address")
|
||||
break
|
||||
|
||||
if public_ip:
|
||||
console.print(f"[dim]Updating SSH config to public IP: {public_ip}[/dim]")
|
||||
try:
|
||||
username = api.get_username()
|
||||
add_ssh_host(
|
||||
config_path=config.ssh.config_path,
|
||||
host_name=ssh_hostname,
|
||||
hostname=public_ip,
|
||||
user=username,
|
||||
identity_file=config.ssh.identity_file,
|
||||
)
|
||||
console.print("[green]✓[/green] SSH config updated to public IP")
|
||||
except Exception as e:
|
||||
if verbose:
|
||||
console.print(f"[dim]Could not update SSH config: {e}[/dim]")
|
||||
if not public_ip:
|
||||
console.print(
|
||||
"[red]Error:[/red] Could not determine public IP for droplet. "
|
||||
"Aborting hibernate to prevent unreachable droplet on wake."
|
||||
)
|
||||
return None
|
||||
|
||||
console.print(f"[dim]Updating SSH config to public IP: {public_ip}[/dim]")
|
||||
try:
|
||||
username = api.get_username()
|
||||
add_ssh_host(
|
||||
config_path=config.ssh.config_path,
|
||||
host_name=ssh_hostname,
|
||||
hostname=public_ip,
|
||||
user=username,
|
||||
identity_file=config.ssh.identity_file,
|
||||
)
|
||||
console.print("[green]✓[/green] SSH config updated to public IP")
|
||||
except (OSError, DigitalOceanAPIError) as e:
|
||||
console.print(
|
||||
f"[red]Error:[/red] Could not update SSH config to public IP: {e}. "
|
||||
"Aborting hibernate to prevent unreachable droplet on wake."
|
||||
)
|
||||
return None
|
||||
|
||||
# Now safe to logout from Tailscale (SSH config points to public IP)
|
||||
console.print("[dim]Logging out from Tailscale...[/dim]")
|
||||
@@ -3805,6 +3817,11 @@ def hibernate(
|
||||
# Note: tailscale_locked was already detected earlier
|
||||
if tailscale_locked:
|
||||
console.print("[dim]Detected Tailscale lockdown[/dim]")
|
||||
console.print(
|
||||
"[yellow]Note:[/yellow] Cannot verify temporary SSH rule on a "
|
||||
"powered-off droplet. If the droplet is unreachable after wake, "
|
||||
"use the DigitalOcean console to access it."
|
||||
)
|
||||
|
||||
_complete_hibernate(
|
||||
api,
|
||||
@@ -3865,7 +3882,11 @@ def hibernate(
|
||||
console.print()
|
||||
|
||||
# Step 0: Prepare for hibernate (handle Tailscale lockdown if present)
|
||||
tailscale_locked = prepare_for_hibernate(config, api, droplet, droplet_name, verbose)
|
||||
prepare_result = prepare_for_hibernate(config, api, droplet, droplet_name, verbose)
|
||||
if prepare_result is None:
|
||||
console.print("[dim]Fix the SSH connectivity issue and retry.[/dim]")
|
||||
raise typer.Exit(1)
|
||||
tailscale_locked = prepare_result
|
||||
|
||||
# Step 1: Power off if not already off
|
||||
if status != "off":
|
||||
@@ -3940,6 +3961,7 @@ def wake(
|
||||
..., autocompletion=complete_snapshot_name, help="Name of the hibernated droplet to restore"
|
||||
),
|
||||
no_tailscale: bool = typer.Option(False, "--no-tailscale", help="Skip Tailscale VPN re-setup"),
|
||||
verbose: bool = typer.Option(False, "--verbose", "-v", help="Show verbose debug output"),
|
||||
):
|
||||
"""
|
||||
Wake a hibernated droplet (restore from snapshot).
|
||||
@@ -4080,7 +4102,7 @@ def wake(
|
||||
identity_file=config.ssh.identity_file,
|
||||
)
|
||||
console.print("[green]✓[/green] SSH config updated")
|
||||
except Exception as e:
|
||||
except OSError as e:
|
||||
console.print(f"[yellow]⚠[/yellow] Could not update SSH config: {e}")
|
||||
|
||||
# Handle Tailscale re-setup if the original droplet had Tailscale lockdown
|
||||
@@ -4107,7 +4129,7 @@ def wake(
|
||||
time.sleep(10)
|
||||
|
||||
# Re-setup Tailscale (clean state from hibernate logout)
|
||||
tailscale_ip = setup_tailscale(ssh_hostname, username, config)
|
||||
tailscale_ip = setup_tailscale(ssh_hostname, username, config, verbose)
|
||||
|
||||
if not tailscale_ip:
|
||||
console.print(
|
||||
|
||||
@@ -6,6 +6,7 @@ from unittest.mock import MagicMock, patch
|
||||
import pytest
|
||||
import typer
|
||||
|
||||
from dropkit.api import DigitalOceanAPIError
|
||||
from dropkit.main import (
|
||||
_resize_hibernated_snapshot,
|
||||
add_temporary_ssh_rule,
|
||||
@@ -332,10 +333,10 @@ class TestPrepareForHibernate:
|
||||
@patch("dropkit.main.tailscale_logout")
|
||||
@patch("dropkit.main.add_temporary_ssh_rule")
|
||||
@patch("dropkit.main.add_ssh_host")
|
||||
def test_temp_rule_failure_skips_logout(
|
||||
def test_temp_rule_failure_aborts(
|
||||
self, mock_add_ssh_host, mock_add_temp_rule, mock_logout, temp_ssh_config
|
||||
):
|
||||
"""Test returns True but skips logout if temp rule fails (safety)."""
|
||||
"""Test returns None to abort hibernate if temp rule fails."""
|
||||
# Create SSH config with Tailscale IP
|
||||
Path(temp_ssh_config).write_text("""Host dropkit.myvm
|
||||
HostName 100.80.123.45
|
||||
@@ -355,11 +356,98 @@ class TestPrepareForHibernate:
|
||||
|
||||
result = prepare_for_hibernate(mock_config, mock_api, mock_droplet, "myvm")
|
||||
|
||||
# Should still return True because we detected Tailscale lockdown
|
||||
assert result is True
|
||||
# But logout should NOT be called (safety - need public IP fallback first)
|
||||
# Should return None to signal failure and abort hibernate
|
||||
assert result is None
|
||||
# Logout should NOT be called
|
||||
mock_logout.assert_not_called()
|
||||
# SSH config should NOT be updated
|
||||
mock_add_ssh_host.assert_not_called()
|
||||
|
||||
@patch("dropkit.main.tailscale_logout")
|
||||
@patch("dropkit.main.add_temporary_ssh_rule")
|
||||
@patch("dropkit.main.add_ssh_host")
|
||||
def test_missing_public_ip_aborts(
|
||||
self, mock_add_ssh_host, mock_add_temp_rule, mock_logout, temp_ssh_config
|
||||
):
|
||||
"""Test returns None to abort hibernate if no public IP found."""
|
||||
Path(temp_ssh_config).write_text("""Host dropkit.myvm
|
||||
HostName 100.80.123.45
|
||||
User ubuntu
|
||||
""")
|
||||
|
||||
mock_config = MagicMock()
|
||||
mock_config.ssh.config_path = temp_ssh_config
|
||||
mock_config.ssh.identity_file = "~/.ssh/id_ed25519"
|
||||
|
||||
mock_api = MagicMock()
|
||||
|
||||
# Droplet with no public IPv4 network
|
||||
mock_droplet = {"networks": {"v4": [{"type": "private", "ip_address": "10.0.0.1"}]}}
|
||||
|
||||
mock_add_temp_rule.return_value = True
|
||||
|
||||
result = prepare_for_hibernate(mock_config, mock_api, mock_droplet, "myvm")
|
||||
|
||||
assert result is None
|
||||
mock_logout.assert_not_called()
|
||||
mock_add_ssh_host.assert_not_called()
|
||||
|
||||
@patch("dropkit.main.tailscale_logout")
|
||||
@patch("dropkit.main.add_temporary_ssh_rule")
|
||||
@patch("dropkit.main.add_ssh_host")
|
||||
def test_ssh_config_update_failure_aborts(
|
||||
self, mock_add_ssh_host, mock_add_temp_rule, mock_logout, temp_ssh_config
|
||||
):
|
||||
"""Test returns None to abort hibernate if SSH config update fails."""
|
||||
Path(temp_ssh_config).write_text("""Host dropkit.myvm
|
||||
HostName 100.80.123.45
|
||||
User ubuntu
|
||||
""")
|
||||
|
||||
mock_config = MagicMock()
|
||||
mock_config.ssh.config_path = temp_ssh_config
|
||||
mock_config.ssh.identity_file = "~/.ssh/id_ed25519"
|
||||
|
||||
mock_api = MagicMock()
|
||||
mock_api.get_username.return_value = "testuser"
|
||||
|
||||
mock_droplet = {"networks": {"v4": [{"type": "public", "ip_address": "203.0.113.50"}]}}
|
||||
|
||||
mock_add_temp_rule.return_value = True
|
||||
mock_add_ssh_host.side_effect = OSError("Permission denied")
|
||||
|
||||
result = prepare_for_hibernate(mock_config, mock_api, mock_droplet, "myvm")
|
||||
|
||||
assert result is None
|
||||
mock_logout.assert_not_called()
|
||||
|
||||
@patch("dropkit.main.tailscale_logout")
|
||||
@patch("dropkit.main.add_temporary_ssh_rule")
|
||||
@patch("dropkit.main.add_ssh_host")
|
||||
def test_ssh_config_api_failure_aborts(
|
||||
self, mock_add_ssh_host, mock_add_temp_rule, mock_logout, temp_ssh_config
|
||||
):
|
||||
"""Test returns None to abort hibernate if API call for username fails."""
|
||||
Path(temp_ssh_config).write_text("""Host dropkit.myvm
|
||||
HostName 100.80.123.45
|
||||
User ubuntu
|
||||
""")
|
||||
|
||||
mock_config = MagicMock()
|
||||
mock_config.ssh.config_path = temp_ssh_config
|
||||
mock_config.ssh.identity_file = "~/.ssh/id_ed25519"
|
||||
|
||||
mock_api = MagicMock()
|
||||
mock_api.get_username.side_effect = DigitalOceanAPIError("API timeout")
|
||||
|
||||
mock_droplet = {"networks": {"v4": [{"type": "public", "ip_address": "203.0.113.50"}]}}
|
||||
|
||||
mock_add_temp_rule.return_value = True
|
||||
|
||||
result = prepare_for_hibernate(mock_config, mock_api, mock_droplet, "myvm")
|
||||
|
||||
assert result is None
|
||||
mock_logout.assert_not_called()
|
||||
# And SSH config should NOT be updated (early return)
|
||||
mock_add_ssh_host.assert_not_called()
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user