# Copyright 2024 - The Android Open Source Project
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""RemoteInstanceDeviceFactory provides basic interface to create a Trusty
device factory."""

import argparse
import json
import logging
import os
import posixpath as remote_path
import shlex
import tempfile
import traceback

from acloud import errors
from acloud.create import create_common
from acloud.internal import constants
from acloud.internal.lib import cvd_utils
from acloud.internal.lib import utils
from acloud.public import report
from acloud.public.actions import gce_device_factory
from acloud.pull import pull


logger = logging.getLogger(__name__)
_CONFIG_JSON_FILENAME = "config.json"

# log files under REMOTE_LOG_FOLDER in order to
# enable `acloud pull` to retrieve them
_REMOTE_LOG_FOLDER = constants.REMOTE_LOG_FOLDER
_REMOTE_STDOUT_PATH = f"{_REMOTE_LOG_FOLDER}/qemu_trusty_console.log"
_REMOTE_STDERR_PATH = f"{_REMOTE_LOG_FOLDER}/qemu_trusty_err.log"

# below Trusty image archive is generated by:
_TRUSTY_IMAGE_PACKAGE = {
    "git_default": "trusty_tee_package_goog.tar.gz",
    "git_main-without-vendor": "trusty_tee_package.tar.gz",
    "trusty_manifest": "trusty_image_package.tar.gz",
}

_DLKM_STAGING = "dlkm_staging"
_KERNEL_STAGING = f"{_DLKM_STAGING}/flatten/lib/modules"
_MODULES_LOAD = "modules.load"

# below Host tools archive is generated by:
# branch: git_main / target: qemu_trusty_arm64
_TRUSTY_HOST_PACKAGE_DIR = "trusty-host_package"
_TRUSTY_HOST_TARBALL = "trusty-host_package.tar.gz"

# Default Trusty build target. This does not depend on the android branch.
_DEFAULT_TRUSTY_BUILD_TARGET = "qemu_generic_arm64_gicv3_test_debug"


def _TrustyImagePackageFilename(build_target, build_branch):
    if build_branch == "trusty_manifest":
        trusty_target = build_target.replace("_", "-")
        return f"{trusty_target}.{_TRUSTY_IMAGE_PACKAGE[build_branch]}"
    if build_branch in _TRUSTY_IMAGE_PACKAGE:
        return _TRUSTY_IMAGE_PACKAGE[build_branch]
    return _TRUSTY_IMAGE_PACKAGE["git_default"]


def _FindHostPackage(package_path=None):
    if package_path:
        # checked in create_args._VerifyTrustyArgs
        return package_path
    dirs_to_check = create_common.GetNonEmptyEnvVars(
        constants.ENV_ANDROID_SOONG_HOST_OUT, constants.ENV_ANDROID_HOST_OUT
    )
    dist_dir = utils.GetDistDir()
    if dist_dir:
        dirs_to_check.append(dist_dir)

    for path in dirs_to_check:
        for name in [_TRUSTY_HOST_TARBALL, _TRUSTY_HOST_PACKAGE_DIR]:
            trusty_host_package = os.path.join(path, name)
            if os.path.exists(trusty_host_package):
                return trusty_host_package
    raise errors.GetTrustyLocalHostPackageError(
        "Can't find the trusty host package (Try lunching a trusty target "
        "like qemu_trusty_arm64-trunk_staging-userdebug and running 'm'): \n"
        + "\n".join(dirs_to_check)
    )


def _FindTrustyImagePackage():
    dist_dir = utils.GetDistDir()
    if dist_dir:
        for name in _TRUSTY_IMAGE_PACKAGE.values():
            trusty_image_package = os.path.join(dist_dir, name)
            if os.path.exists(trusty_image_package):
                return trusty_image_package
    raise errors.GetTrustyLocalImagePackageError(
        "Can't find the trusty image package (Try lunching a trusty target "
        "such as qemu_trusty_arm64-trunk_staging-userdebug and running 'm dist trusty-tee_package')"
    )


class RemoteInstanceDeviceFactory(gce_device_factory.GCEDeviceFactory):
    """A class that can produce a Trusty device."""

    def __init__(self, avd_spec, local_android_image_artifact=None):
        super().__init__(avd_spec, local_android_image_artifact)
        self._all_logs = {}
        self._launch_args = self._ParseLaunchArgs()

    # pylint: disable=broad-except
    def CreateInstance(self):
        """Create and start a single Trusty instance.

        Returns:
            The instance name as a string.
        """
        instance = self.CreateGceInstance()
        if instance in self.GetFailures():
            return instance

        try:
            self._ProcessArtifacts()
            self._StartTrusty()
        except Exception as e:
            self._SetFailures(instance, traceback.format_exception(e))

        self._FindLogFiles(
            instance,
            instance in self.GetFailures() and not self._avd_spec.no_pull_log,
        )
        return instance

    def _SshRun(self, cmd: str, local_substitution=False):
        """run a command in the remote GCE, using single quotes by default
        in order to ensure that substitution happens in the remote shell.
        for example, $PATH or $(pwd) to be substituted in the remote GCE env,
        which is the default desired behavior.

        Some commands (such as input redirection <) however,
        require substitution in the local environment.
        In that case, local_substitution=True shall be used.
        """

        self._ssh.Run(
            cmd if local_substitution else shlex.quote(cmd),
            show_output=True,
            timeout=constants.DEFAULT_SSH_TIMEOUT,
        )

    def _ProcessArtifacts(self):
        """Process artifacts.

        - If images source is local, tool will upload images from local site to
          remote instance.
        - If images source is remote, tool will download images from android
          build to remote instance.
        """
        avd_spec = self._avd_spec
        if avd_spec.image_source == constants.IMAGE_SRC_LOCAL:
            host_package_artifact = _FindHostPackage(avd_spec.trusty_host_package)
            cvd_utils.UploadArtifacts(
                self._ssh,
                cvd_utils.GCE_BASE_DIR,
                (self._local_image_artifact or avd_spec.local_image_dir),
                host_package_artifact,
            )
        elif avd_spec.image_source == constants.IMAGE_SRC_REMOTE:
            self._FetchBuild()
            if not self._TryUseGKIKernelModules():
                # fetch the kernel image from the android build artifacts
                self._FetchAndUploadKernelImage()
        if avd_spec.local_trusty_image:
            self._UploadBuildArchive(avd_spec.local_trusty_image)
        elif avd_spec.image_source == constants.IMAGE_SRC_LOCAL:
            local_trusty_image = _FindTrustyImagePackage()
            self._UploadBuildArchive(local_trusty_image)
        else:
            self._FetchAndUploadTrustyImages()

        config = {
            "linux": "kernel",
            "linux_arch": "arm64",
            "initrd": "ramdisk.img",
            "atf": "atf/qemu/debug",
            "qemu": "bin/trusty_qemu_system_aarch64",
            "extra_qemu_flags": ["-machine", "gic-version=3"],
            "android_image_dir": ".",
            "rpmbd": "bin/rpmb_dev",
            "arch": "arm64",
            "adb": "bin/adb",
        }
        with tempfile.NamedTemporaryFile(
            mode="w+t", suffix=".json"
        ) as config_json_file:
            json.dump(config, config_json_file)
            config_json_file.flush()
            remote_config_path = remote_path.join(
                cvd_utils.GCE_BASE_DIR, _CONFIG_JSON_FILENAME
            )
            self._ssh.ScpPushFile(config_json_file.name, remote_config_path)
            logger.debug(
                "ScpPushFile from %s to %s\n",
                config_json_file.name,
                remote_config_path,
            )

    # We are building our own command-line instead of using
    # self._compute_client.FetchBuild() because we need to use the host cvd
    # tool rather than `fetch_cvd`. The downloaded fetch_cvd tool is too
    # old and cannot handle a custom host package filename. This can be
    # removed when b/298447306 is fixed.
    @utils.TimeExecute(function_description="Fetching builds")
    def _FetchBuild(self):
        """Fetch builds from android build server."""
        avd_spec = self._avd_spec
        build_client = self._compute_client.build_api

        # Provide the default trusty host package artifact filename. We must
        # explicitly use the default build id/branch and target for the host
        # package if those values were not set for the host package so that we
        # can override the artifact filename.
        host_package = avd_spec.host_package_build_info.copy()
        if not (
            host_package[constants.BUILD_ID] or host_package[constants.BUILD_BRANCH]
        ):
            host_package[constants.BUILD_ID] = avd_spec.remote_image[constants.BUILD_ID]
            host_package[constants.BUILD_BRANCH] = avd_spec.remote_image[
                constants.BUILD_BRANCH
            ]
        if not host_package[constants.BUILD_TARGET]:
            host_package[constants.BUILD_TARGET] = avd_spec.remote_image[
                constants.BUILD_TARGET
            ]
        host_package.setdefault(constants.BUILD_ARTIFACT, _TRUSTY_HOST_TARBALL)

        fetch_args = build_client.GetFetchBuildArgs(
            avd_spec.remote_image,
            {},
            avd_spec.kernel_build_info,
            {},
            {},
            {},
            {},
            host_package,
        )
        fetch_cmd = constants.CMD_CVD_FETCH + ["-credential_source=gce"] + fetch_args
        self._SshRun(" ".join(fetch_cmd))

    @utils.TimeExecute(function_description="Fetching & Uploading Trusty image")
    def _FetchAndUploadTrustyImages(self):
        """Fetch Trusty image archive from ab, Upload to GCE"""
        build_client = self._compute_client.build_api
        trusty_build_info = self._avd_spec.trusty_build_info
        if trusty_build_info[constants.BUILD_BRANCH]:
            build_id = trusty_build_info[constants.BUILD_ID]
            build_branch = trusty_build_info[constants.BUILD_BRANCH]
            build_target = (
                trusty_build_info[constants.BUILD_TARGET]
                or _DEFAULT_TRUSTY_BUILD_TARGET
            )
            if not build_id:
                build_id = build_client.GetLKGB(build_target, build_branch)
            trusty_image_package = _TrustyImagePackageFilename(
                build_target, "trusty_manifest"
            )
        else:
            # if Trusty build_branch not specified, use the android build branch
            # get the Trusty image package from the android platform manifest
            android_build_info = self._avd_spec.remote_image
            build_id = android_build_info[constants.BUILD_ID]
            build_branch = android_build_info[constants.BUILD_BRANCH]
            build_target = android_build_info[constants.BUILD_TARGET]
            trusty_image_package = _TrustyImagePackageFilename(None, build_branch)
        with tempfile.NamedTemporaryFile(suffix=".tar.gz") as image_local_file:
            image_local_path = image_local_file.name
            build_client.DownloadArtifact(
                build_target,
                build_id,
                trusty_image_package,
                image_local_path,
            )
            self._UploadBuildArchive(image_local_path)

    @utils.TimeExecute(function_description="Fetching & Uploading Kernel Image")
    def _FetchAndUploadKernelImage(self):
        """Fetch Kernel image from ab, Upload to GCE"""
        build_client = self._compute_client.build_api
        android_build_info = self._avd_spec.remote_image
        build_id = android_build_info[constants.BUILD_ID]
        build_target = android_build_info[constants.BUILD_TARGET]
        with tempfile.NamedTemporaryFile(prefix="kernel") as image_local_file:
            image_local_path = image_local_file.name
            logger.debug('DownloadArtifact "kernel" to %s\n', image_local_path)
            build_client.DownloadArtifact(
                build_target,
                build_id,
                "kernel",
                image_local_path,
            )
            logger.debug("DownloadArtifact (kernel image) to %s\n", image_local_path)
            self._ssh.ScpPushFile(image_local_path, f"{cvd_utils.GCE_BASE_DIR}/kernel")
            logger.debug(
                "ScpPushFile from %s to %s\n",
                image_local_path,
                f"{cvd_utils.GCE_BASE_DIR}/kernel",
            )

    @utils.TimeExecute(function_description="Fetching & Uploading GKI Artifacts")
    def _FetchAndUploadGKIArtifacts(self):
        """Fetch GKI Build artifacts from the kernel and its dynamic module Targets"""
        kernel_build_info = self._avd_spec.kernel_build_info
        build_branch = kernel_build_info[constants.BUILD_BRANCH]
        if not build_branch:
            # if kernel branch is not provided we use the kernel prebuilts in Android
            return False
        kernel_ko_dict = {
            "kernel_aarch64": [
                "virtio_blk.ko",
                "virtio_console.ko",
                "virtio_pci.ko",
                "virtio_pci_modern_dev.ko",
                "virtio_pci_legacy_dev.ko",
            ],
            "kernel_virt_aarch64": [
                "failover.ko",
                "net_failover.ko",
                "virtio_mmio.ko",
                "virtio_net.ko",
                "system_heap.ko",
            ],
            "trusty_aarch64": [
                "ffa-core.ko",
                "ffa-module.ko",
                "trusty-ffa.ko",
                "trusty-smc.ko",
                "trusty-core.ko",
                "trusty-ipc.ko",
                "trusty-log.ko",
                "trusty-test.ko",
                "trusty-virtio.ko",
                "trusty-virtio-polling.ko",
            ],
        }
        build_client = self._compute_client.build_api

        build_id_list = [
            value
            for value in [
                self._launch_args.kernel_trusty_build_id
                or build_client.GetLKGB("trusty_aarch64", build_branch),
                self._launch_args.kernel_virt_build_id
                or build_client.GetLKGB("kernel_virt_aarch64", build_branch),
                kernel_build_info[constants.BUILD_ID]
                or build_client.GetLKGB("kernel_aarch64", build_branch),
            ]
            if value is not None
        ]
        # we use the oldest build_id in the hope that the oldest LKGB
        # has all the necessary targets
        build_id = min(build_id_list)

        def _fetchAndUpload(
            build_target, file_name, dest_dir=None, dest_file_name=None
        ):
            dest_file_path = dest_file_name if dest_file_name else file_name
            dest_file_path = (
                f"{dest_dir}/{dest_file_path}" if dest_dir else dest_file_path
            )
            with tempfile.NamedTemporaryFile(prefix=file_name) as image_local_file:
                image_local_path = image_local_file.name
                build_client.DownloadArtifact(
                    build_target,
                    build_id,
                    file_name,
                    image_local_path,
                )
                self._ssh.ScpPushFile(
                    image_local_path, f"{cvd_utils.GCE_BASE_DIR}/{dest_file_path}"
                )
                logger.debug(
                    "\n_fetchAndUpload %s & ScpPushFile to %s\n",
                    file_name,
                    dest_file_path,
                )

        dlkm_staging_path = remote_path.join(cvd_utils.GCE_BASE_DIR, _DLKM_STAGING)
        self._SshRun(f"mkdir -p {dlkm_staging_path}")

        def _uploadDlkm(build_target):
            with tempfile.NamedTemporaryFile(
                prefix="system_dlkm_staging_archive", suffix=".tar.gz"
            ) as image_local_file:
                image_local_path = image_local_file.name
                build_client.DownloadArtifact(
                    build_target,
                    build_id,
                    "system_dlkm_staging_archive.tar.gz",
                    image_local_path,
                )
                # f"--wildcards ./flatten/lib/modules/*.ko --strip-components=4 "
                self._SshRun(
                    f"tar -xzf - -C {dlkm_staging_path} < {image_local_path}",
                    local_substitution=True,
                )
                # self._SshRun(f"ls -l {dlkm_staging_path}/flatten/lib/modules")
                logger.debug("\n_uploadDlkm done!\n")

        def _uploadModulesLoad():
            modules_load_path = remote_path.join(cvd_utils.GCE_BASE_DIR, _MODULES_LOAD)
            # virtio devices need to be loaded in a determinitic order
            # in order to have a determistic mount (on /dev/vport4p1)
            # of the rpmb block device - see aosp/3473834 for an example on how
            # adding a blk device changes the rpmb device uri
            modules_load = [
                "failover.ko",
                "net_failover.ko",
                "virtio_blk.ko",
                "virtio_console.ko",
                "virtio_mmio.ko",
                "virtio_net.ko",
                "virtio_pci.ko",
                "system_heap.ko",
                "ffa-core.ko",
                "ffa-module.ko",
                "trusty-ffa.ko",
                "trusty-smc.ko",
                "trusty-core.ko",
                "trusty-ipc.ko",
                "trusty-log.ko",
                "trusty-test.ko",
                "trusty-virtio.ko",
                "virtio_pci_modern_dev.ko",
                "virtio_pci_legacy_dev.ko",
            ]
            with tempfile.NamedTemporaryFile(
                prefix=_MODULES_LOAD, mode="w+t"
            ) as modules_load_file:
                modules_load_file.write("\n".join(modules_load))
                modules_load_file.flush()
                self._ssh.ScpPushFile(modules_load_file.name, modules_load_path)
            logger.debug("\n_uploadModulesLoad done!\n")

        _fetchAndUpload("kernel_aarch64", "Image", dest_file_name="kernel")
        _uploadDlkm("kernel_aarch64")
        _uploadModulesLoad()
        for build_target in ["kernel_virt_aarch64", "trusty_aarch64"]:
            for _ko in kernel_ko_dict[build_target]:
                _fetchAndUpload(build_target, _ko, dest_dir=_KERNEL_STAGING)

        # self._SshRun(f"ls -l {kernel_staging_path}")
        return True

    def _TryUseGKIKernelModules(self):
        """Use GKI modules from the kernel build."""

        if not self._FetchAndUploadGKIArtifacts():
            return False
        # remove the inadequate modules files from the staging area
        self._SshRun(f"rm {_KERNEL_STAGING}/modules.*")
        # invoke replace_ramdisk_modules
        # with option --override-modules-load to use
        # the Trusty QEMU specific modules.load
        self._SshRun(
            "PATH=$(pwd)/bin:$PATH ./bin/replace_ramdisk_modules "
            f"--depmod=depmod "
            f"--android-ramdisk=ramdisk.img "
            f"--kernel-ramdisk={_KERNEL_STAGING} "
            f"--output-ramdisk=ramdisk.img "
            f"--override-modules-load {_MODULES_LOAD}"
        )
        # self._SshRun("ls -lst .")
        return True

    def _UploadBuildArchive(self, archive_path):
        self._SshRun(
            f"tar -xzf - -C {cvd_utils.GCE_BASE_DIR} < {archive_path}",
            local_substitution=True,
        )

    def _ParseLaunchArgs(self):
        parser = argparse.ArgumentParser(prog="AVD Launch Args")
        parser.add_argument(
            "--extra-linux-args",
            nargs="+",
            default=[],
            help="Trusty QEMU run.py option\n"
            "allowing to add extra arguments to the linux kernel command line.\n",
        )
        parser.add_argument(
            "--kernel-virt-build-id",
            type=str,
            default=None,
            help="Trusty acloud driver option\n"
            "provide an extra build id for the out-of-band kernel virtual driver modules.\n",
        )
        parser.add_argument(
            "--kernel-trusty-build-id",
            type=str,
            default=None,
            help="Trusty acloud driver option\n"
            "provide an extra build id for the out-of-band kernel Trusty driver modules.\n",
        )
        # see CVD Launch Args
        # exhaustive list at tools/acloud/internal/lib/cvd_utils.py
        # not yet used by Trusty QEMU run.py
        for arg_str in ["data_policy", "config"]:
            parser.add_argument(
                f"-{arg_str}",
                type=str,
                default=None,
                help="CVD specific launch arg\n"
                "not yet used by Trusty QEMU run.py.\n",
            )
        for arg_int in [
            "memory_mb",
            "blank_data_image_mb",
            "x_res",
            "y_res",
            "dpi",
            "cpus",
            "num_AVD",
            "base_instance_num",
        ]:
            parser.add_argument(
                f"-{arg_int}",
                type=int,
                default=None,
                help="CVD specific launch arg\n"
                "not yet used by Trusty QEMU run.py.\n",
            )
        return parser.parse_args(self._avd_spec.launch_args.split())

    @utils.TimeExecute(function_description="Starting Trusty")
    def _StartTrusty(self):
        """Start the model on the GCE instance."""
        self._SshRun(f"mkdir -p {_REMOTE_LOG_FOLDER}")
        # TODO(b/417379600): remove the ln commands when the broken symlink
        # root cause is identified
        self._SshRun("ln -rsf lk.bin atf/qemu/debug/bl32.bin")
        self._SshRun(
            "ln -rsf test-runner/external/trusty/bootloader/test-runner/test-runner.bin atf/qemu/debug/bl33.bin"
        )
        # prepare launch args
        launch_args = (
            f"--extra-linux-args {' '.join(self._launch_args.extra_linux_args)} "
            if len(self._launch_args.extra_linux_args) > 0
            else ""
        )
        # We use an explicit subshell so we can run this command in the
        # background.
        cmd = "-- sh -c " + shlex.quote(
            shlex.quote(
                f"{cvd_utils.GCE_BASE_DIR}/run.py "
                f"--verbose --config={_CONFIG_JSON_FILENAME} "
                f"{launch_args}"
                f"> {_REMOTE_STDOUT_PATH} "
                f"2> {_REMOTE_STDERR_PATH} &"
            )
        )
        self._ssh.Run(cmd, self._avd_spec.boot_timeout_secs or 30, retry=0)

    def _FindLogFiles(self, instance, download):
        """Find and pull all log files from instance.

        Args:
            instance: String, instance name.
            download: Whether to download the files to a temporary directory
                      and show messages to the user.
        """
        logs = [cvd_utils.HOST_KERNEL_LOG]
        if self._avd_spec.image_source == constants.IMAGE_SRC_REMOTE:
            logs.append(cvd_utils.GetRemoteFetcherConfigJson(cvd_utils.GCE_BASE_DIR))
            logs.append(cvd_utils.GetRemoteFetchLog(cvd_utils.GCE_BASE_DIR))
        logs.append(report.LogFile(_REMOTE_STDOUT_PATH, constants.LOG_TYPE_KERNEL_LOG))
        logs.append(report.LogFile(_REMOTE_STDERR_PATH, constants.LOG_TYPE_TEXT))
        self._all_logs[instance] = logs

        logger.debug("logs: %s", logs)
        if download:
            # To avoid long download time, fetch from the first device only.
            log_paths = [log["path"] for log in logs]
            error_log_folder = pull.PullLogs(self._ssh, log_paths, instance)
            self._compute_client.ExtendReportData(
                constants.ERROR_LOG_FOLDER, error_log_folder
            )

    def GetLogs(self):
        """Get all device logs.

        Returns:
            A dictionary that maps instance names to lists of report.LogFile.
        """
        return self._all_logs
