#!/usr/bin/env python3

"""secret-encryption: verify libvirt's encrypted-at-rest secrets feature.

This test checks that:
  1. The secrets key directory is shipped with root:root 0700 permissions.
  2. The virt-secret-init-encryption.service bootstrap unit generates the
     encryption key.
  3. Test encryption of a secret (set and get its value via the API)
"""

import base64
import os
import stat
import subprocess
import sys
import uuid as uuidlib

SECRETS_KEY_DIR = "/var/lib/libvirt/secrets"
ENC_KEY = os.path.join(SECRETS_KEY_DIR, "secrets-encryption-key")
SECRET_CONF = "/etc/libvirt/secret.conf"
# configDir where secret value files (<uuid>.base64) are stored in system mode
SECRET_STORE = "/etc/libvirt/secrets"
CUSTOM_KEY = "/etc/libvirt/custom-secret.key"
SECRET_XML = "/tmp/secret-encryption.xml"


def test_secret_encryption():
    uuid = None

    def run(cmd, **kwargs):
        """Run a command, echoing it first (mirrors the old shell ``set -x``)."""
        print("+ " + " ".join(cmd), file=sys.stderr)
        kwargs.setdefault("check", True)
        return subprocess.run(cmd, **kwargs)

    def run_out(cmd):
        """Run a command and return its stripped stdout as text."""
        print("+ " + " ".join(cmd), file=sys.stderr)
        return subprocess.run(
            cmd, check=True, capture_output=True, text=True
        ).stdout.strip()

    def fail(msg):
        print("ERROR: " + msg, file=sys.stderr)
        sys.exit(1)

    def read_text(path):
        with open(path) as f:
            return f.read()

    def copy_file(src, dst):
        with open(src, "rb") as fsrc, open(dst, "wb") as fdst:
            fdst.write(fsrc.read())

    def append_line(path, line):
        with open(path, "a") as f:
            f.write(line + "\n")

    # The service is pulled in by libvirtd via the 10-secret.conf drop-in; make
    # sure it has run so the encryption key exists (ConditionPathExists makes it
    # a no-op if the key is already present).
    subprocess.run(
        ["systemctl", "reset-failed", "virt-secret-init-encryption.service"],
        stdout=subprocess.DEVNULL,
        stderr=subprocess.DEVNULL,
    )
    run(["systemctl", "start", "virt-secret-init-encryption.service"])
    run(["systemctl", "restart", "libvirtd.service"])

    # 1. Packaging: the key directory must be shipped with root:root 0700 perms
    if not os.path.isdir(SECRETS_KEY_DIR):
        fail(f"{SECRETS_KEY_DIR} is not a directory")
    st = os.stat(SECRETS_KEY_DIR)
    perms = stat.S_IMODE(st.st_mode)
    if st.st_uid != 0 or st.st_gid != 0 or perms != 0o700:
        fail(
            f"{SECRETS_KEY_DIR} must be root:root 0700, "
            f"got uid={st.st_uid} gid={st.st_gid} mode={oct(perms)}"
        )

    # 2. Check that the bootstrap service has produced the key
    if not os.path.isfile(ENC_KEY):
        fail(f"encryption key {ENC_KEY} was not produced")

    # 3. Encrypted-at-rest round trip
    secret_plaintext = "libvirt-secret-encryption-test-value"
    secret_b64 = base64.b64encode(secret_plaintext.encode()).decode()

    # Start from a clean state by removing any pre-existing secrets.
    run(
        [
            "sh",
            "-c",
            "virsh secret-list | awk 'NR>2 && $1 ~ /^[0-9a-f-]/ {print $1}' "
            "| xargs -r -n 1 virsh secret-undefine",
        ]
    )

    secret_volume = f"/var/lib/libvirt/images/secret-encryption-test-{uuidlib.uuid4()}.img"
    with open(SECRET_XML, "w") as f:
        f.write(
            "<secret ephemeral='no' private='no'>\n"
            "  <description>autopkgtest secret encryption</description>\n"
            "  <usage type='volume'>\n"
            f"    <volume>{secret_volume}</volume>\n"
            "  </usage>\n"
            "</secret>\n"
        )

    define_out = run_out(["virsh", "secret-define", SECRET_XML])
    for line in define_out.splitlines():
        if "Secret" in line:
            secret_uuid = line.split()[1]
            break
    if not secret_uuid:
        fail("could not parse generated secret UUID")

    run(["virsh", "secret-set-value", "--secret", secret_uuid, "--base64", secret_b64])

    # The API must still return the original value (driver decrypts transparently)
    if run_out(["virsh", "secret-get-value", secret_uuid]) != secret_b64:
        fail("API did not return the original secret value")

    print("Secret encryption test successful")


if __name__ == "__main__":
    test_secret_encryption()
