#!/usr/bin/env python3
import argparse
import copy
from diagrams import Diagram, Edge, Cluster
from diagrams.aws.enablement import ManagedServices
from diagrams.custom import Custom
import importlib
import json
import os
import traceback
import yaml
from pprint import pprint

# Query YAML data.
def query_path(data, path, default=None):
    paths = path.split(".")
    tmp = data
    for p in paths[:-1]:
        tmp = tmp.get(p)
        if tmp is None:
            return default
    tmp = tmp.get(paths[-1])
    if tmp is None:
        return default
    return tmp

# Directory where this script is.
DIRNAME = os.path.dirname(__file__)

# Load configuration.
config = {}
with open(DIRNAME + "/kube-diagrams.yaml") as f:
    config = yaml.safe_load(f) # load YAML config file

# Get edge config.
def get_edge_config(edge_kind):
    edge_config = config.get("edges", {}).get(edge_kind)
    if edge_config is None:
        print("Warning: %s edge configuration not found!" % edge_kind)
    return edge_config if edge_config else {}

# Get config associated to a resource.
already_warned_node_configs = set()
def get_node_config(resource):
    resource_type = get_type(resource)
    node_config = config.get("nodes", {}).get(resource_type)
    while isinstance(node_config, str):
        if node_config not in already_warned_node_configs:
            print(f"Warning: {resource_type} used instead of {node_config}!")
            already_warned_node_configs.add(node_config)
        node_config = config.get("nodes", {}).get(node_config)
    if node_config is None:
        if resource_type not in already_warned_node_configs:
            print(f"Warning: {resource_type} node configuration undefined!")
            already_warned_node_configs.add(resource_type)
    return node_config if node_config else {}

# Get resource type, i.e., "kind/apiVersion".
def get_type(resource):
    return resource["kind"] + "/" + resource["apiVersion"]

# Get the namespace of a resource.
# If not defined then returns default namespace.
def get_namespace(resource):
    return query_path(resource, "metadata.namespace", config.get("default_namespace", "default"))

# Create diagrams node.
def create_diagram_node(resource, fontcolor="black"):
    # Format node label
    node_label = resource["metadata"]["name"]
    MAX_NODE_LABEL_LENGTH = 16
    if len(node_label) > MAX_NODE_LABEL_LENGTH:
        def split_node_label(node_label, separator):
            parts = node_label.split(separator)
            node_label = parts[0]
            current_length = len(node_label)
            for part in parts[1:]:
                if current_length + len(part) >= MAX_NODE_LABEL_LENGTH:
                    node_label += "\n"
                    current_length = 0
                node_label = node_label + separator + part
                current_length = current_length + 1 + len(part)
            return node_label
        for separator in ["-", "."]:
            if separator in node_label:
                node_label = split_node_label(node_label, separator)
                break
    # Search diagram node class
    diagram_node_class = ManagedServices # default diagram node class
    node_config = get_node_config(resource)
    if node_config != None: # node config found
        custom_icon = node_config.get("custom_icon")
        if custom_icon != None: # custom_icon defined
            return Custom(
                node_label,
                os.path.abspath(custom_icon.replace("$KD", DIRNAME)),
                fontcolor=fontcolor
            )
        # Get diagram node class name
        diagram_node_classname = node_config.get("diagram_node_classname")
        if diagram_node_classname != None: # classname defined
            # Import Diagrams node class module
            idx = diagram_node_classname.rfind('.')
            if idx != -1:
                module = importlib.import_module(diagram_node_classname[:idx])
            # Get diagram node class
            diagram_node_class = getattr(module, diagram_node_classname[idx+1:])
    # Create a diagram node
    return diagram_node_class(node_label, fontcolor=fontcolor)

class ResourceCluster(object):
    def __init__(self):
        self.resources = {} # dict<str, resource>
        self.clusters = {} # dict<str, Cluster>

    def get_or_create_cluster(self, name):
        cluster = self.clusters.get(name)
        if cluster is None:
            cluster = ResourceCluster()
            self.clusters[name] = cluster
        return cluster

class EdgesContext(list):
    def __init__(self, resource):
        list.__init__(self)
        self.resource = resource
        self.namespace = get_namespace(resource)

    def add_resource(self, path):
        target = query_path(self.resource, path)
        self.append([
            "%s/%s/%s/%s" % (
              target["name"],
              self.namespace,
              target["kind"],
              target["apiVersion"],
            ),
            "REFERENCE"
        ])

    def add_resources(self, path, name_path, namespace_path, kind, apiVersion):
        resources = query_path(self.resource, path)
        if resources is None:
            return
        for resource in resources:
            self.append([
                "%s/%s/%s/%s" % (
                    resource[name_path],
                    resource.get(namespace_path, self.namespace),
                    kind,
                    apiVersion
                ),
                "REFERENCE"
            ])

    def add_all_resources_matching_labels(self, kind, path, data=None, resource_labels_path="metadata.labels", edge_kind="SELECTOR"):
        if data is None:
            data = self.resource
        if isinstance(path, str):
            matchLabels = query_path(data, path)
        elif isinstance(path, dict):
            matchLabels = path
        if matchLabels is None:
            return False
        resourceFound = False
        for rid, resource in resources.items():
            if resource["kind"] != kind:
                continue
            labels = query_path(resource, resource_labels_path, {})
            found = True
            for sk, sv in matchLabels.items():
                if labels.get(sk) != sv:
                    found = False
                    break
            if found:
                self.append([rid, edge_kind])
                resourceFound = True
        return resourceFound

    def add_all_volume_resources(self, path):
        def add_volumes(volumes):
            for volume in volumes:
                if "configMap" in volume:
                    configMap = volume["configMap"]
                    configMap_id = f"{configMap.get('name')}/{self.namespace}/ConfigMap/v1"
                    if configMap.get("optional") is True:
                        if configMap_id not in resources:
                            print(f"Warning: optional config map {configMap_id} undefined!")
                            continue # skip it
                    self.append([
                        configMap_id,
                        "REFERENCE"
                    ])
                elif "secret" in volume:
                    secret = volume["secret"]
                    secretName = secret.get("secretName")
                    if secretName is None:
                        # Warning: jenkins chart uses 'name' instead of 'secretName'!
                        secretName = secret.get("name")
                    secret_id = f"{secretName}/{self.namespace}/Secret/v1"
                    if secret.get("optional") is True:
                        if secret_id not in resources:
                            print(f"Warning: optional secret {secret_id} undefined!")
                            continue # skip it
                    self.append([
                        secret_id,
                        "REFERENCE"
                    ])
                elif "persistentVolumeClaim" in volume:
                    self.append([
                        "%s/%s/PersistentVolumeClaim/v1" % (
                            volume["persistentVolumeClaim"]["claimName"],
                            self.namespace
                        ),
                        "REFERENCE"
                    ])
                elif "projected" in volume:
                    add_volumes(query_path(volume, "projected.sources"))

        volumes = query_path(resource, path)
        if volumes != None:
            add_volumes(volumes)
        return

    def add_containers_env_valueFrom(self, path):
        containers = query_path(self.resource, path)
        if containers is None:
            return
        target_resources = set()
        for container in containers:
            for env in query_path(container, "env", []):
                configMapKeyRefName = query_path(env, "valueFrom.configMapKeyRef.name")
                if configMapKeyRefName != None:
                    target_resources.add(
                        "%s/%s/ConfigMap/v1" % (
                            configMapKeyRefName,
                            self.namespace
                        )
                    )
                secretKeyRefName = query_path(env, "valueFrom.secretKeyRef.name")
                if secretKeyRefName != None:
                    secret_id = f"{secretKeyRefName}/{self.namespace}/Secret/v1"
                    if query_path(env, "valueFrom.secretKeyRef.optional") is True:
                        if secret_id not in resources:
                            print(f"Warning: optional secret {secret_id} undefined!")
                            continue # skip it
                    target_resources.add(
                        secret_id
                    )
        for target_resource in target_resources:
            self.append([target_resource, "REFERENCE"])

    def add_volume_claim_templates(self, path):
        volumeClaimTemplates = query_path(self.resource, path)
        if volumeClaimTemplates is None:
            return
        for volumeClaimTemplate in volumeClaimTemplates:
            self.append([
                f"{volumeClaimTemplate['metadata']['name']}/{self.namespace}/PersistentVolumeClaim/v1",
                "REFERENCE"
            ])

    def add_storage_class(self, path, data=None):
        if data is None:
            data = self.resource
        storageClassName = query_path(data, path)
        if storageClassName == "":
            print(f"Warning: {self.resource['kind']} {self.resource['metadata']['name']} {path} equals to \"\"!")
            storageClassName = None
        if storageClassName != None:
            self.append([
                "%s/StorageClass/storage.k8s.io/v1" % storageClassName,
                "REFERENCE"
            ])

    def add_volume(self, path):
        volumeName = query_path(self.resource, path)
        if volumeName != None:
            self.append([
                "%s/PersistentVolume/v1" % volumeName,
                "REFERENCE"
            ])

    def add_wait_for_services(self, path):
        initContainers = query_path(self.resource, path)
        if initContainers is None:
            return
        for ic in initContainers:
            if ic.get("name") == "wait-for-services":
                for arg in query_path(ic, "args", []):
                    if arg.startswith("-service="):
                        sn = arg[len("-service="):]
                        self.append([
                            "%s/%s/Service/v1" % (
                                sn,
                                self.namespace
                            ),
                            "DEPENDENCE"
                        ])

    def add_service_account(self, path):
        podSpec = query_path(self.resource, path)
        serviceAccountName = podSpec.get("serviceAccountName")
        if serviceAccountName is None:
            serviceAccountName = podSpec.get("serviceAccount")
            if serviceAccountName != None:
                print(f"Warning: {self.resource['kind']} {self.resource['metadata']['name']} {path}.serviceAccount deprecated!")
        if serviceAccountName == "":
            print(f"Warning: {self.resource['kind']} {self.resource['metadata']['name']} {path}.serviceAccountName equals to \"\"!")
            serviceAccountName = None
        if serviceAccountName != None:
            self.append([
                "%s/%s/ServiceAccount/v1" % (
                    serviceAccountName,
                    self.namespace
                ),
                "REFERENCE"
            ])

    def add_role(self, path):
        roleRef = self.resource[path]
        apiGroup = query_path(roleRef, "apiGroup", "rbac.authorization.k8s.io")
        if apiGroup == "":
            print(f"Warning: {self.resource['kind']} {self.resource['metadata']['name']} {path} equals to \"\"!")
            apiGroup = "rbac.authorization.k8s.io"
        if roleRef["kind"] == "Role":
            namespace = "/%s" % self.namespace
        else:
            namespace = ""
        self.append([
            "%s%s/%s/%s/v1" % (
                roleRef["name"],
                namespace,
                roleRef["kind"],
                apiGroup
            ),
            "REFERENCE"
      ])

    def add_subjects(self):
        for subject in query_path(self.resource, "subjects", []):
            namespace = subject.get("namespace")
            if namespace is None and subject["kind"] == "ServiceAccount":
                namespace = get_namespace(self.resource)
            if namespace != None:
                apiVersion = query_path(subject, "apiGroup", "v1")
                if apiVersion == "":
                    apiVersion = "v1"
                rid = "%s/%s/%s/%s" % (
                    subject["name"],
                    namespace,
                    subject["kind"],
                    apiVersion
                )
            else:
                rid = "%s/%s/%s/v1" % (
                    subject["name"],
                    subject["kind"],
                    query_path(subject, "apiGroup", "rbac.authorization.k8s.io")
                )
            self.append([rid,"REFERENCE-UP"])

    def add_service(self, path, data=None):
        if data is None:
            data = self.resource
        name = query_path(data, path)
        if name == "":
            print(f"Warning: {self.resource['kind']} {self.resource['metadata']['name']} {path} equals to \"\"!")
            name = None
        if name != None:
            self.append([
                "%s/%s/Service/v1" % (
                    name,
                    self.namespace
                ),
                "REFERENCE"
            ])

    def add_owned_resources(self):
        uid = query_path(self.resource, "metadata.uid")
        for _, resource in resources.items():
            for ownerReference in query_path(resource, "metadata.ownerReferences", []):
                if ownerReference.get("uid") == uid:
                    self.append([
                        f"{resource['metadata']['name']}/{resource['metadata']['namespace']}/{resource['kind']}/{resource['apiVersion']}",
                        "OWNER"
                    ])

    def add_ingress_and_egress_rules(self):
        ingress_rules = query_path(self.resource, "spec.ingress", [])
        for ingress_rule in ingress_rules:
            for ingress_from in query_path(ingress_rule, "from", []):
                podSelector = query_path(ingress_from, "podSelector")
                self.add_all_resources_matching_labels("Pod", "matchLabels", podSelector, edge_kind="SELECTOR-INGRESS")
        #TODO: support for other forms of rules

    def add_webhooks(self):
        for webhook in query_path(self.resource, "webhooks", []):
            service = query_path(webhook, "clientConfig.service")
            if service != None:
                self.append([
                    "%s/%s/Service/v1" % (
                        service["name"],
                        service["namespace"]
                    ),
                    "REFERENCE"
                ])

    def add_all_workload_resources(self, path, custom_workload_resources={}, edge_kind="SELECTOR"):
        selector = query_path(self.resource, path)
        if selector is None:
            return
        build_in_workload_resources = {
            "Pod": "metadata.labels",
            "Deployment": "spec.template.metadata.labels",
            "ReplicaSet": "spec.template.metadata.labels",
            "StatefulSet": "spec.template.metadata.labels",
            "DaemonSet": "spec.template.metadata.labels",
            "Job": "spec.template.metadata.labels",
        }
        for kind, label in {**build_in_workload_resources, **custom_workload_resources}.items():
            resourceNotFound = True
            if self.add_all_resources_matching_labels(
                    kind,
                    path,
                    resource_labels_path=label,
                    edge_kind=edge_kind):
                resourceNotFound = False
                break
        if resourceNotFound:
            resource_name = self.resource["metadata"]["name"]
            selector = query_path(self.resource, path)
            print(f"Warning: No workload resource matches labels({selector}) for Service({resource_name})!")

    def add_edges_for_service(self, custom_workload_resources={}):
        self.add_all_workload_resources("spec.selector", custom_workload_resources)
        self.add_all_resources_matching_labels(
                "EndpointSlice",
                { "kubernetes.io/service-name": query_path(self.resource, "metadata.name") }
        )
        rid = "%s/%s/%s/%s" % (
            query_path(self.resource, "metadata.name"),
            get_namespace(self.resource),
            "Endpoints",
            "v1"
        )
        if rid in resources:
            self.append([rid, "SELECTOR"])

    def add_networks(self, path):
      annotations = query_path(resource, "spec.template.metadata.annotations")
      if annotations != None:
        networks = annotations.get("k8s.v1.cni.cncf.io/networks")
        if networks != None:
          for network in json.loads(networks):
            resource_name = query_path(resource, "metadata.name")
            edges.append([
              "%s/%s/NetworkAttachmentDefinition/k8s.cni.cncf.io/v1" % (
                network["name"],
                get_namespace(resource)
              ),
              "REFERENCE"
            ])

    def add_priority_class(self, path):
      priorityClassName = query_path(self.resource, path)
      if priorityClassName == "":
          print(f"Warning: {self.resource['kind']} {self.resource['metadata']['name']} {path} equals to \"\"!")
          priorityClassName = None
      if priorityClassName != None:
        self.append([
            f"{priorityClassName}/PriorityClass/scheduling.k8s.io/v1",
            "REFERENCE"
        ])

# Parse arguments
parser = argparse.ArgumentParser(
    prog="kube-diagrams",
    description="Generate Kubernetes architecture diagrams from Kubernetes manifest files")
parser.add_argument("filename", nargs='+',
    help="the Kubernetes manifest filename to process")
parser.add_argument("-o", "--output", type=str,
    help="output diagram filename")
parser.add_argument("-f", "--format", type=str,
    help="output format, allowed formats are png (default), jpg, svg, pdf, and dot",
    default="png")
parser.add_argument("-c", "--config", type=str,
    help="custom kube-diagrams configuration file")
parser.add_argument("-v", "--verbose",
    help="verbosity, set to false by default",
    action="store_true", default=False)
parser.add_argument("--without-namespace",
    help="disable namespace cluster generation",
    action="store_true", default=False)
args = parser.parse_args()

# Process arguments.
if args.output is None:
    args.output = args.filename[0][:args.filename[0].rfind('.')]

if args.config != None:
    with open(args.config) as f:
        custom_config = yaml.safe_load(f) # load custom config file
        if custom_config: # not empty file
            if "default_namespace" in custom_config:
                config["default_namespace"] = custom_config["default_namespace"]
            if custom_config.get("edges"):
                config["edges"].update(custom_config["edges"])
            if custom_config.get("clusters"):
                config["clusters"].extend(custom_config["clusters"])
            if custom_config.get("nodes"):
                config["nodes"].update(custom_config["nodes"])

# Load the Kubernetes manifest file.
resources = {} # a map of all resources
resource_cluster = ResourceCluster() # Clustering resources
for filename in args.filename:
    print("Load %s..." % filename)
    with open(filename) as f:
        def process_resource(resource):
            cluster = resource_cluster
            name = resource["metadata"]["name"]
            if get_node_config(resource).get("scope") == "Cluster":
                rid = name + "/" + get_type(resource)
                resources[rid] = resource
                if not(args.without_namespace) and resource["kind"] == "Namespace":
                    cluster = resource_cluster.get_or_create_cluster("Namespace: %s" % name)
            else: # scope = Namespaced
                rid = name + "/" + get_namespace(resource) + "/" + get_type(resource)
                resources[rid] = resource
                if not args.without_namespace:
                    cluster = resource_cluster.get_or_create_cluster("Namespace: %s" % get_namespace(resource))

            def process_clusters(cluster, resource, clusters):
                for cluster_config in clusters:
                    if "label" in cluster_config:
                        label = cluster_config["label"]
                        labels = query_path(resource, "metadata.labels")
                        if labels != None and label in labels:
                            cluster = cluster.get_or_create_cluster("%s: %s" % (
                                cluster_config["title"],
                                labels[label]
                            ))
                            clusters.remove(cluster_config)
                            return process_clusters(cluster, resource, clusters)
                return cluster

            cluster = process_clusters(
                cluster,
                resource,
                copy.deepcopy(config.get("clusters", []))
            )
            cluster.resources[rid] = resource
            code_to_exec = get_node_config(resource).get("nodes")
            if code_to_exec != None:
                nodes = []
                try:
                    exec(code_to_exec)
                except Exception as exc:
                    print("Error:", type(exc), ":", exc.args)
                    traceback.print_exception(exc)
                    print("Edges script:\n", code_to_exec)
                    print("Resource:")
                    pprint(resource)
                    raise
                for node in nodes:
                    process_resource(node)

        data = yaml.safe_load_all(f) # load YAML file
        for resource in data:
            if resource is None:
                continue # Skip empty YAML content
            if resource["kind"] == "List":
                for r in resource["items"]:
                    process_resource(r)
            else:
                process_resource(resource)

# Print loaded Kubernetes resources
if args.verbose:
    def process_cluster(cluster, ident=0):
        for rid, resource in cluster.resources.items():
            print("  " * ident, "- Resource %s" % rid)
        for cid, sub_cluster in cluster.clusters.items():
            print("  " * ident, "- %s" % cid)
            process_cluster(sub_cluster, ident + 1)
    print("Loaded Kubernetes resources:")
    process_cluster(resource_cluster)

# Generate diagram
with Diagram("", filename=args.output, show=False, direction="TB", outformat=args.format):
    # Generate diagram nodes
    diagram_nodes = {}
    def create_nodes(cluster):
        for rid, resource in cluster.resources.items():
            diagram_nodes[rid] = create_diagram_node(resource)
        for cid, sub_cluster in cluster.clusters.items():
            with Cluster(cid):
                create_nodes(sub_cluster)
    create_nodes(resource_cluster)

    # Generate diagram edges
    undefined_nodes = {}
    UNDEFINED_EDGE = query_path(config, "edges.UNDEFINED", {})
    for resource_id, resource in resources.items():
        edges = EdgesContext(resource)
        # execute config code to generate edges
        code_to_exec = get_node_config(resource).get("edges")
        if code_to_exec:
            try:
                exec(code_to_exec)
            except Exception as exc:
                print("Error:", type(exc), ":", exc.args)
                traceback.print_exception(exc)
                print("Edges script:\n", code_to_exec)
                print("Resource:")
                pprint(resource)
                raise
        # generate diagram edges
        for edge in edges:
            edge_rid = edge[0]
            edge_name = edge[1]
            try:
                edge_config = get_edge_config(edge_name)
                if edge_config.get("direction") == "up":
                    diagram_nodes[edge_rid] >> Edge(color="white", style="invisible") >> diagram_nodes[resource_id]
                diagram_nodes[resource_id] >> Edge(**edge_config) >> diagram_nodes[edge_rid]
            except KeyError as ke:
                if edge_rid in config["cluster-resources"]:
                    print(f"Info: referenced {edge_rid} resource is provided by K8s clusters.")
                    continue # skip this edge as the resource is provided by K8s clusters.
                print("Error: %s resource not found!" % ke)
                tmp = str(edge_rid).split("/")
                if tmp[2] in ["ServiceAccount", "Service", "ConfigMap", "Secret", "PersistentVolumeClaim", "Role", "Issuer", "Deployment"]: #TODO: avoid specific cases
                    rsrc = {
                        "apiVersion": "/".join(tmp[3:]),
                        "kind": tmp[2],
                        "metadata": {
                            "name": tmp[1] + "/" + tmp[0]
                        }
                    }
                else:
                    rsrc = {
                        "apiVersion": tmp[2] + "/" + tmp[3],
                        "kind": tmp[1],
                        "metadata": {
                            "name": tmp[0]
                        }
                    }
                diagram_node = undefined_nodes.get(edge_rid)
                if diagram_node is None:
                    diagram_node = create_diagram_node(rsrc, UNDEFINED_EDGE.get("color"))
                    undefined_nodes[edge_rid] = diagram_node
                diagram_nodes[resource_id] >> Edge(**UNDEFINED_EDGE) >> diagram_node

print("%s.%s generated." % (args.output, args.format))
