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
26 changes: 26 additions & 0 deletions api/core/network.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
import ipaddress
import socket


def is_internal_address(hostname: str) -> bool:
"""Return True if the hostname is, or resolves to, an internal network
address: loopback, RFC 1918 private, link-local, reserved, or multicast.
Unresolvable hostnames are not considered internal."""
try:
ips = [ipaddress.ip_address(hostname)]
except ValueError:
# hostname is a name rather than a literal IP — resolve it.
try:
results = socket.getaddrinfo(hostname, None, socket.AF_UNSPEC)
ips = [ipaddress.ip_address(str(r[4][0]).split("%")[0]) for r in results]
except socket.gaierror:
return False

return any(
ip.is_loopback
or ip.is_private
or ip.is_link_local
or ip.is_reserved
or ip.is_multicast
for ip in ips
)
2 changes: 1 addition & 1 deletion api/experimentation/warehouse_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from rest_framework import serializers

from core.network import is_internal_address
from experimentation.models import WarehouseConnection, WarehouseType
from experimentation.types import (
CLICKHOUSE_DEFAULTS,
Expand All @@ -10,7 +11,6 @@
ClickHouseCredentials,
SnowflakeConfig,
)
from webhooks.fields import is_internal_address


def validate_clickhouse_credentials(
Expand Down
9 changes: 1 addition & 8 deletions api/tests/unit/core/test_fields.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,19 +19,12 @@ def test_get_prep_value__json_value__returns_ciphertext_that_roundtrips() -> Non
assert field.from_db_value(stored, None, None) == value


def test_get_prep_value__none__returns_none() -> None:
def test_field_methods__none__returns_none() -> None:
# Given
field = EncryptedJSONField()

# When & Then
assert field.get_prep_value(None) is None


def test_from_db_value__none__returns_none() -> None:
# Given
field = EncryptedJSONField()

# When & Then
assert field.from_db_value(None, None, None) is None


Expand Down
69 changes: 20 additions & 49 deletions api/tests/unit/experimentation/test_services.py
Original file line number Diff line number Diff line change
Expand Up @@ -2060,14 +2060,31 @@ def test_verify_clickhouse_connection__reachable__sets_connected(
} in log.events


def test_verify_clickhouse_connection__driver_error__sets_errored(
@pytest.mark.parametrize(
"credentials, execute_side_effect, expected_detail",
[
(
{"password": "hunter2"},
Exception("connection refused"),
"Verification failed.",
),
(None, None, "Stored connection details are incomplete."),
],
ids=["driver_error", "missing_credentials"],
)
def test_verify_clickhouse_connection__failure__sets_errored_with_detail(
clickhouse_connection: WarehouseConnection,
credentials: dict[str, str] | None,
execute_side_effect: Exception | None,
expected_detail: str,
log: StructuredLogCapture,
mocker: MockerFixture,
) -> None:
# Given
mock_client = mocker.patch("experimentation.services.Client")
mock_client.return_value.execute.side_effect = Exception("connection refused")
mock_client.return_value.execute.side_effect = execute_side_effect
clickhouse_connection.credentials = credentials
clickhouse_connection.save()
failure_count_before = _verification_count("failure")

# When
Expand All @@ -2076,59 +2093,13 @@ def test_verify_clickhouse_connection__driver_error__sets_errored(
# Then
clickhouse_connection.refresh_from_db()
assert clickhouse_connection.status == WarehouseConnectionStatus.ERRORED
assert clickhouse_connection.status_detail == "Verification failed."
mock_client.return_value.disconnect.assert_called_once_with()
assert clickhouse_connection.status_detail == expected_detail
assert _verification_count("failure") == failure_count_before + 1
assert any(
event["event"] == "connection.verification_failed" for event in log.events
)


def test_verify_clickhouse_connection__missing_credentials__sets_errored(
clickhouse_connection: WarehouseConnection,
mocker: MockerFixture,
) -> None:
# Given
mocker.patch("experimentation.services.Client")
clickhouse_connection.credentials = None
clickhouse_connection.save()

# When
verify_clickhouse_connection(clickhouse_connection)

# Then
clickhouse_connection.refresh_from_db()
assert clickhouse_connection.status == WarehouseConnectionStatus.ERRORED
assert (
clickhouse_connection.status_detail
== "Stored connection details are incomplete."
)


def test_verify_clickhouse_connection__environment_lookup_fails__sets_errored(
clickhouse_connection: WarehouseConnection,
mocker: MockerFixture,
) -> None:
# Given
mocker.patch("experimentation.services.Client")
mocker.patch.object(
Environment,
"project",
new_callable=mocker.PropertyMock,
side_effect=Exception("database unavailable"),
)
failure_count_before = _verification_count("failure")

# When
verify_clickhouse_connection(clickhouse_connection)

# Then
clickhouse_connection.refresh_from_db()
assert clickhouse_connection.status == WarehouseConnectionStatus.ERRORED
assert clickhouse_connection.status_detail == "Verification failed."
assert _verification_count("failure") == failure_count_before + 1


@pytest.mark.parametrize(
"error,expected_detail",
[
Expand Down
Loading
Loading