Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 25 additions & 14 deletions opgee/XMLFile.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
'''
"""
.. Created as part of pygcam (2015)
Imported into opgee (2021)

.. Copyright (c) 2015-2022 Richard Plevin
See the https://opensource.org/licenses/MIT for license details.
'''
"""

from io import BytesIO

from lxml import etree as ET
Expand All @@ -15,12 +16,20 @@

_logger = getLogger(__name__)

class XMLFile(object):

parsed_schemas = {} # cache parsed schemas to avoid re-reading and parsing opgee.xsd
class XMLFile:
parsed_schemas = {} # cache parsed schemas to avoid re-reading and parsing opgee.xsd

def __init__(self, filename, xml_string=None, load=True, schemaPath=None,
removeComments=True, conditionalXML=False, varDict=None):
def __init__(
self,
filename,
xml_string=None,
load=True,
schemaPath=None,
removeComments=True,
conditionalXML=False,
varDict=None,
):
"""
Stores information about an XML file; provides wrapper to parse and access
the file tree, and handle "conditional XML".
Expand All @@ -41,26 +50,26 @@ def __init__(self, filename, xml_string=None, load=True, schemaPath=None,
self.xml_string = str.encode(xml_string) if xml_string else None
self.tree = None
self.conditionalXML = conditionalXML
self.varDict = varDict or getConfigDict(section=getParam('OPGEE.DefaultProject'))
self.varDict = varDict or getConfigDict(section=getParam("OPGEE.DefaultProject"))
self.removeComments = removeComments

self.schemaPath = schemaPath
self.schemaPath = schemaPath
self.schemaStream = None

# if filename and load:
if load:
self.read()

def getRoot(self):
'Return the root node of the parse tree'
"Return the root node of the parse tree"
return self.tree.getroot()

def getTree(self):
'Return XML parse tree.'
"Return XML parse tree."
return self.tree

def getFilename(self):
'Return the filename for this ``XMLFile``'
"Return the filename for this ``XMLFile``"
return self.filename

def read(self):
Expand All @@ -86,7 +95,7 @@ def read(self):
raise XmlFormatError(f"Can't read from XML {thing}: {e}")

if self.removeComments:
for elt in tree.iterfind('.//comment'):
for elt in tree.iterfind(".//comment"):
parent = elt.getparent()
if parent is not None:
parent.remove(elt)
Expand Down Expand Up @@ -114,7 +123,7 @@ def validate(self, raiseOnError=True):
# use the cached version if available
schema = self.parsed_schemas.get(self.schemaPath)
if not schema:
ref = imp.files('opgee') / self.schemaPath
ref = imp.files("opgee") / self.schemaPath

with imp.as_file(ref) as path:
xsd = ET.parse(path)
Expand All @@ -126,7 +135,9 @@ def validate(self, raiseOnError=True):
schema.assertValid(tree)
return True
except ET.DocumentInvalid as e:
raise XmlFormatError(f"Validation of '{self.filename}'\n using schema '{self.schemaPath}' failed:\n {e}")
raise XmlFormatError(
f"Validation of '{self.filename}'\n using schema '{self.schemaPath}' failed:\n {e}"
)
else:
valid = schema.validate(tree)
return valid
2 changes: 1 addition & 1 deletion opgee/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,4 +2,4 @@

# warnings.filterwarnings("ignore", category=DeprecationWarning)
# warnings.filterwarnings("error", category=UserWarning) # turn warning into error to debug
warnings.filterwarnings("ignore", category=UserWarning) # turn warning into error to debug
warnings.filterwarnings("ignore", category=UserWarning) # turn warning into error to debug
24 changes: 11 additions & 13 deletions opgee/analysis.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from .config import getParamAsList
from .container import Container
from .core import elt_name, OpgeeObject
from .core import OpgeeObject, elt_name
from .emissions import Emissions
from .error import OpgeeException
from .field import Field
Expand All @@ -20,7 +20,7 @@

class Group(OpgeeObject):
def __init__(self, elt):
self.is_regex = getBooleanXML(elt.attrib.get('regex', 0))
self.is_regex = getBooleanXML(elt.attrib.get("regex", 0))
self.text = elt.text


Expand All @@ -39,6 +39,7 @@ class Analysis(Container):

See also :doc:`OPGEE XML documentation <opgee-xml>`
"""

def __init__(self, name, parent=None, attr_dict=None, field_names=None, groups=None):
super().__init__(name, attr_dict=attr_dict, parent=parent)
self.check_attr_constraints(self.attr_dict)
Expand All @@ -49,22 +50,22 @@ def __init__(self, name, parent=None, attr_dict=None, field_names=None, groups=N
self.model = model = parent

# self.field_dict = None
self._field_names = field_names # may be extended in add_children()
self._field_names = field_names # may be extended in add_children()
self.groups = [] if groups is None else groups

self.fn_unit = self.attr("functional_unit")
self.boundary = self.attr("boundary")

# Create validation sets from system.cfg to avoid hard-coding these
self.functional_units = set(getParamAsList('OPGEE.FunctionalUnits'))
self.functional_units = set(getParamAsList("OPGEE.FunctionalUnits"))

# This is set in use_GWP() below to a pandas Series holding the current
# values in use, indexed by gas name.
self.gwp = None

# Use the GWP years and version specified in XML
gwp_horizon = self.attr('GWP_horizon')
gwp_version = self.attr('GWP_version')
gwp_horizon = self.attr("GWP_horizon")
gwp_version = self.attr("GWP_version")

self.use_GWP(gwp_horizon, gwp_version)

Expand All @@ -75,8 +76,7 @@ def __init__(self, name, parent=None, attr_dict=None, field_names=None, groups=N
text = group.text
if group.is_regex:
prog = re.compile(text)
matches = [field for field in model.fields() for
name in field.group_names if prog.match(name)]
matches = [field for field in model.fields() for name in field.group_names if prog.match(name)]
else:
matches = [field for field in model.fields() if text in field.group_names]

Expand All @@ -97,7 +97,6 @@ def restrict_fields(self, field_names):
# Use list comprehension rather than set.intersection to maintain original order
self._field_names = [name for name in self._field_names if name in names]


def get_field(self, name, raiseError=True) -> Field:
"""
Find a `Field` by name in an `Analysis`.
Expand Down Expand Up @@ -125,8 +124,7 @@ def field_names(self, enabled_only=True):
if enabled_only:
names = [f.name for f in self.fields()]
return names
else:
return self._field_names
return self._field_names

def first_field(self):
return self.get_field(self._field_names[0])
Expand Down Expand Up @@ -227,8 +225,8 @@ def from_xml(cls, elt, parent=None, field_names=None):
"""
name = elt_name(elt)
attr_dict = cls.instantiate_attrs(elt)
field_names = field_names or [elt_name(node) for node in elt.findall('FieldRef')]
groups = [Group(node) for node in elt.findall('Group')]
field_names = field_names or [elt_name(node) for node in elt.findall("FieldRef")]
groups = [Group(node) for node in elt.findall("Group")]

obj = Analysis(name, attr_dict=attr_dict, parent=parent, field_names=field_names, groups=groups)
return obj
Loading
Loading