aboutsummaryrefslogtreecommitdiffhomepage
diff options
context:
space:
mode:
authorAdam Janovsky2021-11-02 10:57:11 +0100
committerAdam Janovsky2021-11-02 10:57:11 +0100
commit5b0d100e6663eb6c0e8ed395ba6774aeada3f779 (patch)
treec3537d0f436da9a16703673d017cd170e0c41d73
parente6c0469072718b8c69ee618a665d72d67bb0ee0f (diff)
downloadsec-certs-5b0d100e6663eb6c0e8ed395ba6774aeada3f779.tar.gz
sec-certs-5b0d100e6663eb6c0e8ed395ba6774aeada3f779.tar.zst
sec-certs-5b0d100e6663eb6c0e8ed395ba6774aeada3f779.zip
improve serialization
-rw-r--r--sec_certs/dataset/common_criteria.py6
-rw-r--r--sec_certs/dataset/cpe.py4
-rw-r--r--sec_certs/dataset/dataset.py21
-rw-r--r--sec_certs/dataset/fips.py6
-rw-r--r--sec_certs/serialization.py7
-rw-r--r--tests/test_cc_heuristics.py2
6 files changed, 16 insertions, 30 deletions
diff --git a/sec_certs/dataset/common_criteria.py b/sec_certs/dataset/common_criteria.py
index a6ea7a61..9c7e6b37 100644
--- a/sec_certs/dataset/common_criteria.py
+++ b/sec_certs/dataset/common_criteria.py
@@ -169,12 +169,6 @@ class CCDataset(Dataset, ComplexSerializableType):
return [(x, self.web_dir / y) for y, x in self.csv_products.items() if 'archived' in y]
@classmethod
- def from_json(cls, input_path: Union[str, Path]):
- dset = Dataset.from_json(input_path)
- dset.set_local_paths()
- return dset
-
- @classmethod
def from_web_latest(cls):
with tempfile.TemporaryDirectory() as tmp_dir:
dset_path = Path(tmp_dir) / 'cc_latest_dataset.json'
diff --git a/sec_certs/dataset/cpe.py b/sec_certs/dataset/cpe.py
index 516eaaf4..355b2dec 100644
--- a/sec_certs/dataset/cpe.py
+++ b/sec_certs/dataset/cpe.py
@@ -25,7 +25,7 @@ logger = logging.getLogger(__name__)
@dataclass
class CPEDataset(ComplexSerializableType):
was_enhanced_with_vuln_cpes: bool
- _json_path: Path
+ json_path: Path
cpes: Dict[str, CPE]
vendor_to_versions: Dict[str, Set[str]] = field(init=False) # Look-up dict cpe_vendor: list of viable versions
vendor_version_to_cpe: Dict[Tuple[str, str], Set[CPE]] = field(init=False) # Look-up dict (cpe_vendor, cpe_version): List of viable cpe items
@@ -102,7 +102,7 @@ class CPEDataset(ComplexSerializableType):
@classmethod
def from_json(cls, input_path: Union[str, Path]):
dset = ComplexSerializableType.from_json(input_path)
- dset._json_path = input_path
+ dset.json_path = input_path
return dset
@classmethod
diff --git a/sec_certs/dataset/dataset.py b/sec_certs/dataset/dataset.py
index c41510d8..d242f6ec 100644
--- a/sec_certs/dataset/dataset.py
+++ b/sec_certs/dataset/dataset.py
@@ -15,7 +15,7 @@ import sec_certs.constants as constants
import sec_certs.parallel_processing as cert_processing
from sec_certs.sample.certificate import Certificate
-from sec_certs.serialization import CustomJSONDecoder, CustomJSONEncoder
+from sec_certs.serialization import CustomJSONDecoder, CustomJSONEncoder, ComplexSerializableType
from sec_certs.config.configuration import config
from sec_certs.serialization import serialize
from sec_certs.dataset.cpe import CPEDataset
@@ -25,7 +25,7 @@ from sec_certs.model.cpe_matching import CPEClassifier
logger = logging.getLogger(__name__)
-class Dataset(ABC):
+class Dataset(ABC, ComplexSerializableType):
def __init__(self, certs: Dict[str, 'Certificate'], root_dir: Path, name: str = 'dataset name',
description: str = 'dataset_description'):
self._root_dir = root_dir
@@ -102,21 +102,16 @@ class Dataset(ABC):
f'The actual number of certs in dataset ({len(dset)}) does not match the claimed number ({claimed}).')
return dset
- def to_json(self, output_path: Union[str, Path] = None):
- if not output_path:
- output_path = self.json_path
-
- with Path(output_path).open('w') as handle:
- json.dump(self, handle, indent=4, cls=CustomJSONEncoder, ensure_ascii=False)
-
@classmethod
def from_json(cls, input_path: Union[str, Path]):
- input_path = Path(input_path)
- with input_path.open('r') as handle:
- dset = json.load(handle, cls=CustomJSONDecoder)
- dset.root_dir = input_path.parent.absolute()
+ dset = ComplexSerializableType.from_json(input_path)
+ dset.root_dir = Path(input_path).parent.absolute()
+ dset.set_local_paths()
return dset
+ def set_local_paths(self):
+ raise NotImplementedError('Not meant to be implemented by the base class.')
+
@abstractmethod
def get_certs_from_web(self):
raise NotImplementedError('Not meant to be implemented by the base class.')
diff --git a/sec_certs/dataset/fips.py b/sec_certs/dataset/fips.py
index a8b0cbcf..46e6b590 100644
--- a/sec_certs/dataset/fips.py
+++ b/sec_certs/dataset/fips.py
@@ -213,12 +213,6 @@ class FIPSDataset(Dataset, ComplexSerializableType):
dset.finalize_results()
return dset
- @classmethod
- def from_json(cls, input_path: Union[str, Path]):
- dset = super().from_json(input_path)
- dset.set_local_paths()
- return dset
-
def set_local_paths(self):
cert: FIPSCertificate
for cert in self.certs.values():
diff --git a/sec_certs/serialization.py b/sec_certs/serialization.py
index f4e56116..8760f8c2 100644
--- a/sec_certs/serialization.py
+++ b/sec_certs/serialization.py
@@ -1,7 +1,7 @@
import json
from datetime import date
from pathlib import Path
-from typing import Dict, List, Union
+from typing import Dict, List, Union, Optional
import copy
@@ -25,8 +25,10 @@ class ComplexSerializableType:
except TypeError as e:
raise TypeError(f'Dict: {dct} on {cls.__mro__}') from e
- def to_json(self, output_path: Union[str, Path] = None):
+ def to_json(self, output_path: Optional[Union[str, Path]] = None):
if not output_path:
+ if not hasattr(self, 'json_path'):
+ raise ValueError(f'The object {self} of type {self.__class__} does not have json_path attribute but to_json() was called without an argument.')
output_path = self.json_path
with Path(output_path).open('w') as handle:
json.dump(self, handle, indent=4, cls=CustomJSONEncoder, ensure_ascii=False)
@@ -38,6 +40,7 @@ class ComplexSerializableType:
obj = json.load(handle, cls=CustomJSONDecoder)
return obj
+
# Decorator for serialization
def serialize(func: callable):
def inner_func(*args, **kwargs):
diff --git a/tests/test_cc_heuristics.py b/tests/test_cc_heuristics.py
index b2e439e6..05b07712 100644
--- a/tests/test_cc_heuristics.py
+++ b/tests/test_cc_heuristics.py
@@ -68,7 +68,7 @@ class TestCommonCriteriaHeuristics(TestCase):
def test_load_cpe_dataset(self):
json_cpe_dset = CPEDataset.from_json(self.data_dir_path / 'auxillary_datasets' / 'cpe_dataset.json')
- json_cpe_dset._json_path = Path('../')
+ json_cpe_dset.json_path = Path('../')
self.assertEqual(self.cpe_dset, json_cpe_dset, 'CPE template dataset does not match CPE dataset loaded from json.')
def test_cpe_lookup_dicts(self):