From 38a0cfebc8bfe449d76649f73260dbb33546dfa3 Mon Sep 17 00:00:00 2001 From: illgamhoduck Date: Fri, 17 Jan 2020 13:08:37 -0800 Subject: [PATCH 1/2] BaseNode / Edge class text input protection added (#423) BaseNode class - Add protection to name, axis_names *Protected in 3 place *Initialize stage - __init__ *Function use setting - set_name / add_axis_names *Property - Add @property to name to protect direct adding node.name = 123 Edge class - Add protection to name *Protected in 3 place *Initialize stage - __init__ *Function use setting - set_name *Property --- tensornetwork/network_components.py | 36 ++++++++++++++++++++++++++++- 1 file changed, 35 insertions(+), 1 deletion(-) diff --git a/tensornetwork/network_components.py b/tensornetwork/network_components.py index 38f51a995..8dc31d176 100644 --- a/tensornetwork/network_components.py +++ b/tensornetwork/network_components.py @@ -69,10 +69,18 @@ def __init__(self, """ self.is_disabled = False - self.name = name if name is not None else '__unnamed_node__' + if not name: + name = '__unnamed_node__' + else: + if type(name) != str: + raise TypeError("Node name should be str type") + self.name = name self.backend = backend self._shape = shape if axis_names is not None: + for axis_name in axis_names: + if type(axis_name) != str: + raise TypeError("Axis_names should be str type") self._edges = [ Edge(node1=self, axis1=i, name=edge_name) for i, edge_name in enumerate(axis_names) @@ -125,6 +133,9 @@ def add_axis_names(self, axis_names: List[Text]) -> None: raise ValueError("axis_names is not the same length as the tensor shape." "axis_names length: {}, tensor.shape length: {}".format( len(axis_names), len(self.shape))) + for axis_name in axis_names: + if type(axis_name) != str: + raise TypeError("Axis_names should be str type") self.axis_names = axis_names[:] def add_edge(self, @@ -312,6 +323,8 @@ def get_all_dangling(self) -> Set["Edge"]: return {edge for edge in self.edges if edge.is_dangling()} def set_name(self, name) -> None: + if type(name) != str: + raise TypeError("Node name should be str type") self.name = name def has_nondangling_edge(self) -> bool: @@ -373,6 +386,16 @@ def edges(self, edges: List) -> None: self.name)) self._edges = edges + @property + def name(self) -> Text: + return self._name + + @name.setter + def name(self, name) -> None: + if type(name) != str: + raise TypeError("Node name should be str type") + self._name = name + @property def axis_names(self) -> List[Text]: return self._axis_names @@ -382,8 +405,12 @@ def axis_names(self, axis_names: List[Text]) -> None: if len(axis_names) != len(self.shape): raise ValueError("Expected {} names, only got {}.".format( len(self.shape), len(axis_names))) + for axis_name in axis_names: + if type(axis_name) != str: + raise TypeError("Axis_names should be str type") self._axis_names = axis_names + @property def signature(self) -> Optional[int]: if self.is_disabled: @@ -810,6 +837,9 @@ def __init__(self, self.is_disabled = False if not name: name = '__unnamed_edge__' + else: + if type(name) != str: + raise TypeError("Edge name should be str type") self._name = name self.node1 = node1 self._axis1 = axis1 @@ -844,6 +874,8 @@ def name(self, name) -> None: if self.is_disabled: raise ValueError( 'Edge has been disabled, setting its name is no longer possible') + if type(name) != str: + raise TypeError("Edge name should be str type") self._name = name @property @@ -988,6 +1020,8 @@ def is_being_used(self) -> bool: return result def set_name(self, name: Text) -> None: + if type(name) != str: + raise TypeError("Edge name should be str type") self.name = name def _save_edge(self, edge_group: h5py.Group) -> None: From 05d3850d59c3ee505024b4073829734a349b09da Mon Sep 17 00:00:00 2001 From: Hyunbyung Park Date: Fri, 17 Jan 2020 19:31:42 -0800 Subject: [PATCH 2/2] BaseNode / Edge class text input protection code revise (#423) *if type(name) != str *if not isinstance(name, str) *change using type to isinstance to follow pylint --- tensornetwork/network_components.py | 24 ++++++++++++------------ 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/tensornetwork/network_components.py b/tensornetwork/network_components.py index 8dc31d176..b6569872f 100644 --- a/tensornetwork/network_components.py +++ b/tensornetwork/network_components.py @@ -72,15 +72,15 @@ def __init__(self, if not name: name = '__unnamed_node__' else: - if type(name) != str: + if not isinstance(name, str): raise TypeError("Node name should be str type") self.name = name self.backend = backend self._shape = shape if axis_names is not None: for axis_name in axis_names: - if type(axis_name) != str: - raise TypeError("Axis_names should be str type") + if not isinstance(axis_name, str): + raise TypeError("axis_names should be str type") self._edges = [ Edge(node1=self, axis1=i, name=edge_name) for i, edge_name in enumerate(axis_names) @@ -134,8 +134,8 @@ def add_axis_names(self, axis_names: List[Text]) -> None: "axis_names length: {}, tensor.shape length: {}".format( len(axis_names), len(self.shape))) for axis_name in axis_names: - if type(axis_name) != str: - raise TypeError("Axis_names should be str type") + if not isinstance(axis_name, str): + raise TypeError("axis_names should be str type") self.axis_names = axis_names[:] def add_edge(self, @@ -323,7 +323,7 @@ def get_all_dangling(self) -> Set["Edge"]: return {edge for edge in self.edges if edge.is_dangling()} def set_name(self, name) -> None: - if type(name) != str: + if not isinstance(name, str): raise TypeError("Node name should be str type") self.name = name @@ -392,7 +392,7 @@ def name(self) -> Text: @name.setter def name(self, name) -> None: - if type(name) != str: + if not isinstance(name, str): raise TypeError("Node name should be str type") self._name = name @@ -406,8 +406,8 @@ def axis_names(self, axis_names: List[Text]) -> None: raise ValueError("Expected {} names, only got {}.".format( len(self.shape), len(axis_names))) for axis_name in axis_names: - if type(axis_name) != str: - raise TypeError("Axis_names should be str type") + if not isinstance(axis_name, str): + raise TypeError("axis_names should be str type") self._axis_names = axis_names @@ -838,7 +838,7 @@ def __init__(self, if not name: name = '__unnamed_edge__' else: - if type(name) != str: + if not isinstance(name, str): raise TypeError("Edge name should be str type") self._name = name self.node1 = node1 @@ -874,7 +874,7 @@ def name(self, name) -> None: if self.is_disabled: raise ValueError( 'Edge has been disabled, setting its name is no longer possible') - if type(name) != str: + if not isinstance(name, str): raise TypeError("Edge name should be str type") self._name = name @@ -1020,7 +1020,7 @@ def is_being_used(self) -> bool: return result def set_name(self, name: Text) -> None: - if type(name) != str: + if not isinstance(name, str): raise TypeError("Edge name should be str type") self.name = name