99
1010import dataclasses
1111import datetime
12+ import functools
1213import logging
1314import os
1415import shutil
4647logger = logging .getLogger (__name__ )
4748
4849
49- @dataclasses .dataclass (frozen = True )
50+ @dataclasses .dataclass (eq = True )
51+ @functools .total_ordering
5052class VulnerabilitySeverity :
5153 # FIXME: this should be named scoring_system, like in the model
5254 system : ScoringSystem
@@ -55,25 +57,23 @@ class VulnerabilitySeverity:
5557 published_at : Optional [datetime .datetime ] = None
5658
5759 def to_dict (self ):
58- published_at_dict = (
59- {"published_at" : self .published_at .isoformat ()} if self .published_at else {}
60- )
61- return {
60+ data = {
6261 "system" : self .system .identifier ,
6362 "value" : self .value ,
6463 "scoring_elements" : self .scoring_elements ,
65- ** published_at_dict ,
6664 }
67-
68- def __eq__ (self , other ):
69- if not isinstance (other , VulnerabilitySeverity ):
70- return NotImplemented
71- return str (self .to_dict ()) == str (other .to_dict ())
65+ if self .published_at :
66+ data ["published_at" ] = self .published_at .isoformat ()
67+ return data
7268
7369 def __lt__ (self , other ):
7470 if not isinstance (other , VulnerabilitySeverity ):
7571 return NotImplemented
76- return str (self .to_dict ()) < str (other .to_dict ())
72+ return self ._cmp_key () < other ._cmp_key ()
73+
74+ # TODO: Add cache
75+ def _cmp_key (self ):
76+ return (str (self .system ), self .value , self .scoring_elements , self .published_at )
7777
7878 @classmethod
7979 def from_dict (cls , severity : dict ):
@@ -89,7 +89,8 @@ def from_dict(cls, severity: dict):
8989 )
9090
9191
92- @dataclasses .dataclass (frozen = True )
92+ @dataclasses .dataclass (eq = True )
93+ @functools .total_ordering
9394class Reference :
9495 reference_id : str = ""
9596 reference_type : str = ""
@@ -100,31 +101,22 @@ def __post_init__(self):
100101 if not self .url :
101102 raise TypeError ("Reference must have a url" )
102103
103- def normalized (self ):
104- severities = sorted (self .severities )
105- return Reference (
106- reference_id = self .reference_id ,
107- url = self .url ,
108- severities = severities ,
109- reference_type = self .reference_type ,
110- )
111-
112- def __eq__ (self , other ):
113- if not isinstance (other , Reference ):
114- return NotImplemented
115- return str (self .to_dict ()) == str (other .to_dict ())
116-
117104 def __lt__ (self , other ):
118105 if not isinstance (other , Reference ):
119106 return NotImplemented
120- return str (self .to_dict ()) < str (other .to_dict ())
107+ return self ._cmp_key () < other ._cmp_key ()
108+
109+ # TODO: Add cache
110+ def _cmp_key (self ):
111+ return (self .reference_id , self .reference_type , self .url , tuple (self .severities ))
121112
122113 def to_dict (self ):
114+ """Return a normalized dictionary representation"""
123115 return {
124116 "reference_id" : self .reference_id ,
125117 "reference_type" : self .reference_type ,
126118 "url" : self .url ,
127- "severities" : [severity .to_dict () for severity in self .severities ],
119+ "severities" : [severity .to_dict () for severity in sorted ( self .severities ) ],
128120 }
129121
130122 @classmethod
@@ -160,7 +152,8 @@ class NoAffectedPackages(Exception):
160152 """
161153
162154
163- @dataclasses .dataclass (frozen = True )
155+ @functools .total_ordering
156+ @dataclasses .dataclass (eq = True )
164157class AffectedPackage :
165158 """
166159 Relate a Package URL with a range of affected versions and a fixed version.
@@ -190,15 +183,14 @@ def get_fixed_purl(self):
190183 raise ValueError (f"Affected Package { self .package !r} does not have a fixed version" )
191184 return update_purl_version (purl = self .package , version = str (self .fixed_version ))
192185
193- def __eq__ (self , other ):
194- if not isinstance (other , AffectedPackage ):
195- return NotImplemented
196- return str (self .to_dict ()) == str (other .to_dict ())
197-
198186 def __lt__ (self , other ):
199187 if not isinstance (other , AffectedPackage ):
200188 return NotImplemented
201- return str (self .to_dict ()) < str (other .to_dict ())
189+ return self ._cmp_key () < other ._cmp_key ()
190+
191+ # TODO: Add cache
192+ def _cmp_key (self ):
193+ return (str (self .package ), str (self .affected_version_range ), str (self .fixed_version ))
202194
203195 @classmethod
204196 def merge (
@@ -304,7 +296,6 @@ class AdvisoryData:
304296 date_published : Optional [datetime .datetime ] = None
305297 weaknesses : List [int ] = dataclasses .field (default_factory = list )
306298 url : Optional [str ] = None
307- created_by : Optional [str ] = None
308299
309300 def __post_init__ (self ):
310301 if self .date_published and not self .date_published .tzinfo :
0 commit comments