Files
stack/tests/api/test_provision.py
kert 2c0a69a4d8 add HKDF credential derivation and auto-rotation
Single root key derives all 18 service credentials via HKDF-SHA256.
Bootstrap/service tiers rotate on key change vs every commit.
Provisioners update PostgreSQL roles and Gitea tokens automatically.
CI pipeline runs `api.auth provision` after each deploy.
2026-02-28 22:45:14 -05:00

192 lines
6.3 KiB
Python

"""Tests for credential provisioning."""
from __future__ import annotations
from unittest.mock import patch
from api.auth.manifest import (
CREDENTIALS,
MANAGED_VARS,
POSTGRES_DATABASES,
POSTGRES_ROLES,
Format,
Provisioner,
Tier,
)
from api.auth.provision import derive_all, write_env
ROOT = bytes.fromhex("deadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeef")
COMMIT = "abc1234"
class TestManifestInvariants:
def test_no_duplicate_env_vars(self):
env_vars = [c.env_var for c in CREDENTIALS]
assert len(env_vars) == len(set(env_vars))
def test_all_managed_vars_in_credentials(self):
assert MANAGED_VARS == frozenset(c.env_var for c in CREDENTIALS)
def test_postgres_roles_have_credentials(self):
for env_var in POSTGRES_ROLES.values():
assert env_var in MANAGED_VARS
def test_postgres_roles_have_databases(self):
for role in POSTGRES_ROLES:
assert role in POSTGRES_DATABASES
def test_aliases_share_scope(self):
by_var = {c.env_var: c for c in CREDENTIALS}
for cred in CREDENTIALS:
if cred.alias_of:
parent = by_var[cred.alias_of]
assert cred.scope == parent.scope
assert cred.fmt == parent.fmt
def test_credential_count(self):
assert len(CREDENTIALS) == 20
class TestDeriveAll:
def test_returns_dict(self):
values = derive_all(ROOT, COMMIT)
assert isinstance(values, dict)
def test_skips_skip_provisioner(self):
values = derive_all(ROOT, COMMIT)
skip_vars = {
c.env_var for c in CREDENTIALS if c.provisioner is Provisioner.SKIP
}
for var in skip_vars:
assert var not in values
def test_hex_format_values(self):
values = derive_all(ROOT, COMMIT)
for cred in CREDENTIALS:
if cred.provisioner is Provisioner.SKIP:
continue
if cred.fmt is Format.HEX:
v = values[cred.env_var]
assert len(v) == 64
bytes.fromhex(v)
def test_password_format_values(self):
values = derive_all(ROOT, COMMIT)
for cred in CREDENTIALS:
if cred.provisioner is Provisioner.SKIP:
continue
if cred.fmt is Format.PASSWORD:
v = values[cred.env_var]
assert "+" not in v
assert "/" not in v
def test_deterministic(self):
a = derive_all(ROOT, COMMIT)
b = derive_all(ROOT, COMMIT)
assert a == b
def test_different_commit_different_service_values(self):
a = derive_all(ROOT, "commit-a")
b = derive_all(ROOT, "commit-b")
# Bootstrap values should be the same
for cred in CREDENTIALS:
if cred.tier is Tier.BOOTSTRAP and cred.provisioner is not Provisioner.SKIP:
assert a[cred.env_var] == b[cred.env_var]
# At least one service value should differ
service_vars = [
c.env_var
for c in CREDENTIALS
if c.tier is Tier.SERVICE
and c.provisioner is not Provisioner.SKIP
and c.alias_of is None
]
assert any(a[v] != b[v] for v in service_vars)
def test_aliases_match_parent(self):
values = derive_all(ROOT, COMMIT)
for cred in CREDENTIALS:
if cred.alias_of and cred.provisioner is not Provisioner.SKIP:
assert values[cred.env_var] == values[cred.alias_of]
class TestWriteEnv:
def test_creates_new_file(self, tmp_path):
p = tmp_path / ".env"
write_env({"A": "1", "B": "2"}, p)
lines = p.read_text().splitlines()
assert lines == ["A=1", "B=2"]
def test_preserves_unmanaged_keys(self, tmp_path):
p = tmp_path / ".env"
p.write_text("DATABRICKS_TOKEN=ext\nOLD_KEY=val\n")
write_env({"NEW_KEY": "new"}, p)
content = p.read_text()
assert "DATABRICKS_TOKEN=ext" in content
assert "NEW_KEY=new" in content
assert "OLD_KEY=val" in content
def test_overwrites_managed_keys(self, tmp_path):
p = tmp_path / ".env"
p.write_text("MY_KEY=old\n")
write_env({"MY_KEY": "new"}, p)
lines = p.read_text().splitlines()
assert "MY_KEY=new" in lines
assert "MY_KEY=old" not in lines
def test_sorted_output(self, tmp_path):
p = tmp_path / ".env"
write_env({"Z": "3", "A": "1", "M": "2"}, p)
lines = p.read_text().splitlines()
keys = [l.split("=")[0] for l in lines]
assert keys == sorted(keys)
def test_skips_comments_and_blanks(self, tmp_path):
p = tmp_path / ".env"
p.write_text("# comment\n\nKEY=val\n")
write_env({"NEW": "v"}, p)
content = p.read_text()
assert "KEY=val" in content
assert "NEW=v" in content
# Comments are not preserved (by design)
assert "# comment" not in content
class TestProvisionPostgres:
def test_calls_docker_exec(self):
values = derive_all(ROOT, COMMIT)
with patch("api.auth.provision.subprocess.run") as mock_run:
from api.auth.provision import provision_postgres
provision_postgres(values, container="test-pg")
assert mock_run.call_count == len(POSTGRES_ROLES)
for call in mock_run.call_args_list:
cmd = call[0][0]
assert cmd[:3] == ["docker", "exec", "test-pg"]
assert "ALTER ROLE" in cmd[-1]
class TestProvisionGitea:
def test_calls_change_password(self):
values = derive_all(ROOT, COMMIT)
mock_client = type("C", (), {})()
mock_client.get = lambda *a, **kw: type("R", (), {"json": lambda self: []})()
mock_client.post = lambda *a, **kw: type(
"R", (), {"json": lambda self: {"sha1": "tok123"}}
)()
with (
patch("api.auth.provision.subprocess.run") as mock_run,
patch(
"api.auth.provision._make_gitea_client",
return_value=mock_client,
),
):
from api.auth.provision import provision_gitea
token = provision_gitea(values, container="test-gitea")
assert token == "tok123"
assert mock_run.call_count == 1
cmd = mock_run.call_args[0][0]
assert "change-password" in cmd