aboutsummaryrefslogtreecommitdiffhomepage
path: root/sec_certs/model/dependency_vulnerability_finder.py
blob: dd25774d2d0fed6072f0e24f2f130e746be89482 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
import logging
from dataclasses import dataclass, field
from enum import Enum
from typing import Dict, List, Optional, Set

from sec_certs.sample.certificate import Certificate
from sec_certs.serialization.json import ComplexSerializableType


class DependencyType(Enum):
    DIRECT = "direct"
    INDIRECT = "indirect"


@dataclass
class DependencyCVE(ComplexSerializableType):
    direct_dependency_cves: Optional[Set[str]] = field(default=None)
    indirect_dependency_cves: Optional[Set[str]] = field(default=None)


Certificates = Dict[str, Certificate]
Vulnerabilities = Dict[str, Dict[str, Optional[Set[str]]]]


class DependencyVulnerabilityFinder:
    """
    The class assigns vulnerabilities to each certificate instance caused by dependencies among certificate instances.
    Adheres to sklearn BaseEstimator interface.
    """

    def __init__(self):
        self.vulnerabilities: Vulnerabilities = {}
        self.certificates: Certificates = {}

    def _overwrite_previous_state(self, certificates: Certificates) -> None:
        self.vulnerabilities = {}
        self.certificates = certificates

    def _get_dataset_cert_ids_occurrences(self) -> Dict[str, int]:
        cert_id_occurrences: Dict[str, int] = {}

        for dgst in self.certificates:
            cert_id = self.certificates[dgst].heuristics.cert_id

            if cert_id is None:
                continue

            cert_id_occurrences[cert_id] = cert_id_occurrences.get(cert_id, 0) + 1

        return cert_id_occurrences

    def _get_cert_dependency_cves(self, dgst: str, dependency_type: DependencyType) -> Optional[Set[str]]:
        dependency_type_dict = {
            DependencyType.DIRECT: self.certificates[dgst].heuristics.report_references.directly_referenced_by,
            DependencyType.INDIRECT: self.certificates[dgst].heuristics.report_references.indirectly_referenced_by,
        }

        references = dependency_type_dict[dependency_type]

        if not references:
            return None

        vulnerabilities = set()
        dataset_cert_id_occurrences = self._get_dataset_cert_ids_occurrences()

        for cert_id in references:
            if cert_id is None:
                continue

            cert_id_occurrences = dataset_cert_id_occurrences.get(cert_id)

            if cert_id_occurrences is None or cert_id_occurrences >= 2:
                continue

            for dgst in self.certificates:
                cert_obj = self.certificates[dgst]
                cert_obj_cves = cert_obj.heuristics.related_cves

                if cert_obj.heuristics.cert_id == cert_id and cert_obj_cves:
                    vulnerabilities.update(cert_obj_cves)

        return vulnerabilities if vulnerabilities else None

    def fit(self, certificates: Certificates) -> Vulnerabilities:
        """
        Method assigns each certificate vulnerabilities caused by dependencies among certificates

        :param Certificates certificates: Dictionary of certificates with digests
        :return Vulnerabilities: Dictionary of vulnerabilities of certificate instances
        """
        self._overwrite_previous_state(certificates)

        cert_id_occurrences = self._get_dataset_cert_ids_occurrences()
        thrown_away_cert_counter = 0

        for dgst in self.certificates:
            cert_id = self.certificates[dgst].heuristics.cert_id

            if cert_id is None:
                continue

            if cert_id_occurrences[cert_id] >= 2:
                thrown_away_cert_counter += 1
                continue

            self.vulnerabilities[dgst] = {}
            self.vulnerabilities[dgst][DependencyType.DIRECT.value] = self._get_cert_dependency_cves(
                dgst, DependencyType.DIRECT
            )
            self.vulnerabilities[dgst][DependencyType.INDIRECT.value] = self._get_cert_dependency_cves(
                dgst, DependencyType.INDIRECT
            )

        if thrown_away_cert_counter > 0:
            logging.warning("There were total of %s certificates skipped due to duplicity", thrown_away_cert_counter)

        return self.vulnerabilities

    def predict_single_cert(self, dgst: str) -> DependencyCVE:
        """
        Method returns vulnerabilities for certificate digest

        :param str dgst: Digest of certificate
        :return DependencyCVE: DependencyCVE object of certificate
        """
        if not self.vulnerabilities.get(dgst):
            return DependencyCVE(direct_dependency_cves=None, indirect_dependency_cves=None)

        return DependencyCVE(
            self.vulnerabilities[dgst][DependencyType.DIRECT.value],
            self.vulnerabilities[dgst][DependencyType.INDIRECT.value],
        )

    def predict(self, dgst_list: List[str]) -> Dict[str, DependencyCVE]:
        """
        Method returns vulnerabilities for a list of certificate digests

        :param List[str] dgst_list: list of certificate digests
        :return Dict[str, DependencyCVE]: Dictionary of DependencyCVE objects for specified certificate digests
        """
        cert_vulnerabilities = {}

        for dgst in dgst_list:
            cert_vulnerabilities[dgst] = self.predict_single_cert(dgst)

        return cert_vulnerabilities