Refactor box validation when wrapping

These became functions only at the behest of pylint.  I don't really
agree with that justification anymore, and flake8 is perfectly happy
with them as methods, so back to methods they go.
This commit is contained in:
jevans 2015-03-19 18:50:00 -04:00
commit d7a3c4df27

View file

@ -840,7 +840,7 @@ class Jp2k(Jp2kBox):
if boxes is None: if boxes is None:
boxes = self._get_default_jp2_boxes() boxes = self._get_default_jp2_boxes()
_validate_jp2_box_sequence(boxes) self._validate_jp2_box_sequence(boxes)
with open(filename, 'wb') as ofile: with open(filename, 'wb') as ofile:
for box in boxes: for box in boxes:
@ -1486,7 +1486,7 @@ class Jp2k(Jp2kBox):
dtypes, nrows, ncols = [], [], [] dtypes, nrows, ncols = [], [], []
for k in range(raw_image.contents.numcomps): for k in range(raw_image.contents.numcomps):
component = raw_image.contents.comps[k] component = raw_image.contents.comps[k]
dtypes.append(_component2dtype(component)) dtypes.append(self._component2dtype(component))
nrows.append(component.h) nrows.append(component.h)
ncols.append(component.w) ncols.append(component.w)
is_cube = all(r == nrows[0] and c == ncols[0] and d == dtypes[0] is_cube = all(r == nrows[0] and c == ncols[0] and d == dtypes[0]
@ -1500,7 +1500,7 @@ class Jp2k(Jp2kBox):
for k in range(raw_image.contents.numcomps): for k in range(raw_image.contents.numcomps):
component = raw_image.contents.comps[k] component = raw_image.contents.comps[k]
_validate_nonzero_image_size(nrows[k], ncols[k], k) self._validate_nonzero_image_size(nrows[k], ncols[k], k)
addr = ctypes.addressof(component.data.contents) addr = ctypes.addressof(component.data.contents)
with warnings.catch_warnings(): with warnings.catch_warnings():
@ -1517,6 +1517,38 @@ class Jp2k(Jp2kBox):
return image return image
def _component2dtype(self, component):
"""Determin the appropriate numpy datatype for an OpenJPEG component.
Parameters
----------
component : ctypes pointer to ImageCompType (image_comp_t)
single image component structure.
Returns
-------
dtype : builtins.type
numpy datatype to be used to construct an image array.
"""
if component.sgnd:
if component.prec <= 8:
dtype = np.int8
elif component.prec <= 16:
dtype = np.int16
else:
msg = "Unhandled precision: {0} bits.".format(component.prec)
raise RuntimeError(msg)
else:
if component.prec <= 8:
dtype = np.uint8
elif component.prec <= 16:
dtype = np.uint16
else:
msg = "Unhandled precision: {0} bits.".format(component.prec)
raise RuntimeError(msg)
return dtype
def get_codestream(self, header_only=True): def get_codestream(self, header_only=True):
"""Returns a codestream object. """Returns a codestream object.
@ -1645,76 +1677,41 @@ class Jp2k(Jp2kBox):
self._comptparms = comptparms self._comptparms = comptparms
def _validate_nonzero_image_size(self, nrows, ncols, component_index):
def _component2dtype(component):
"""Take an OpenJPEG component structure and determine the numpy datatype.
Parameters
----------
component : ctypes pointer to ImageCompType (image_comp_t)
single image component structure.
Returns
-------
dtype : builtins.type
numpy datatype to be used to construct an image array.
"""
if component.sgnd:
if component.prec <= 8:
dtype = np.int8
elif component.prec <= 16:
dtype = np.int16
else:
raise RuntimeError("Unhandled precision, datatype")
else:
if component.prec <= 8:
dtype = np.uint8
elif component.prec <= 16:
dtype = np.uint16
else:
raise RuntimeError("Unhandled precision, datatype")
return dtype
def _validate_nonzero_image_size(nrows, ncols, component_index):
"""The image cannot have area of zero. """The image cannot have area of zero.
""" """
if nrows == 0 or ncols == 0: if nrows == 0 or ncols == 0:
# Letting this situation continue would segfault Python. # Letting this situation continue would segfault openjpeg.
msg = "Component {0} has dimensions {1} x {2}" msg = "Component {0} has dimensions {1} x {2}"
msg = msg.format(component_index, nrows, ncols) msg = msg.format(component_index, nrows, ncols)
raise IOError(msg) raise IOError(msg)
def _validate_jp2_box_sequence(self, boxes):
JP2_IDS = ['colr', 'cdef', 'cmap', 'jp2c', 'ftyp', 'ihdr', 'jp2h', 'jP ',
'pclr', 'res ', 'resc', 'resd', 'xml ', 'ulst', 'uinf', 'url ',
'uuid']
def _validate_jp2_box_sequence(boxes):
"""Run through series of tests for JP2 box legality. """Run through series of tests for JP2 box legality.
This is non-exhaustive. This is non-exhaustive.
""" """
_validate_signature_compatibility(boxes) JP2_IDS = ['colr', 'cdef', 'cmap', 'jp2c', 'ftyp', 'ihdr', 'jp2h',
_validate_jp2h(boxes) 'jP ', 'pclr', 'res ', 'resc', 'resd', 'xml ', 'ulst',
_validate_jp2c(boxes) 'uinf', 'url ', 'uuid']
self._validate_signature_compatibility(boxes)
self._validate_jp2h(boxes)
self._validate_jp2c(boxes)
if boxes[1].brand == 'jpx ': if boxes[1].brand == 'jpx ':
_validate_jpx_box_sequence(boxes) self._validate_jpx_box_sequence(boxes)
else: else:
# Validate the JP2 box IDs. # Validate the JP2 box IDs.
count = _collect_box_count(boxes) count = self._collect_box_count(boxes)
for box_id in count.keys(): for box_id in count.keys():
if box_id not in JP2_IDS: if box_id not in JP2_IDS:
msg = "The presence of a '{0}' box requires that the file " msg = "The presence of a '{0}' box requires that the file "
msg += "type brand be set to 'jpx '." msg += "type brand be set to 'jpx '."
raise IOError(msg.format(box_id)) raise IOError(msg.format(box_id))
_validate_jp2_colr(boxes) self._validate_jp2_colr(boxes)
def _validate_jp2_colr(self, boxes):
def _validate_jp2_colr(boxes):
""" """
Validate JP2 requirements on colour specification boxes. Validate JP2 requirements on colour specification boxes.
""" """
@ -1722,20 +1719,19 @@ def _validate_jp2_colr(boxes):
jp2h = lst[0] jp2h = lst[0]
for colr in [box for box in jp2h.box if box.box_id == 'colr']: for colr in [box for box in jp2h.box if box.box_id == 'colr']:
if colr.approximation != 0: if colr.approximation != 0:
msg = "A JP2 colr box cannot have a non-zero approximation field." msg = "A JP2 colr box cannot have a non-zero approximation "
msg += "field."
raise IOError(msg) raise IOError(msg)
def _validate_jpx_box_sequence(self, boxes):
def _validate_jpx_box_sequence(boxes):
"""Run through series of tests for JPX box legality.""" """Run through series of tests for JPX box legality."""
_validate_label(boxes) self._validate_label(boxes)
_validate_jpx_brand(boxes, boxes[1].brand) self._validate_jpx_brand(boxes, boxes[1].brand)
_validate_jpx_compatibility(boxes, boxes[1].compatibility_list) self._validate_jpx_compatibility(boxes, boxes[1].compatibility_list)
_validate_singletons(boxes) self._validate_singletons(boxes)
_validate_top_level(boxes) self._validate_top_level(boxes)
def _validate_signature_compatibility(self, boxes):
def _validate_signature_compatibility(boxes):
"""Validate the file signature and compatibility status.""" """Validate the file signature and compatibility status."""
# Check for a bad sequence of boxes. # Check for a bad sequence of boxes.
# 1st two boxes must be 'jP ' and 'ftyp' # 1st two boxes must be 'jP ' and 'ftyp'
@ -1749,8 +1745,7 @@ def _validate_signature_compatibility(boxes):
msg = "The ftyp box must contain 'jp2 ' in the compatibility list." msg = "The ftyp box must contain 'jp2 ' in the compatibility list."
raise IOError(msg) raise IOError(msg)
def _validate_jp2c(self, boxes):
def _validate_jp2c(boxes):
"""Validate the codestream box in relation to other boxes.""" """Validate the codestream box in relation to other boxes."""
# jp2c must be preceeded by jp2h # jp2c must be preceeded by jp2h
jp2h_lst = [idx for (idx, box) in enumerate(boxes) jp2h_lst = [idx for (idx, box) in enumerate(boxes)
@ -1768,10 +1763,9 @@ def _validate_jp2c(boxes):
msg = "The codestream box must be preceeded by a jp2 header box." msg = "The codestream box must be preceeded by a jp2 header box."
raise IOError(msg) raise IOError(msg)
def _validate_jp2h(self, boxes):
def _validate_jp2h(boxes):
"""Validate the JP2 Header box.""" """Validate the JP2 Header box."""
_check_jp2h_child_boxes(boxes, 'top-level') self._check_jp2h_child_boxes(boxes, 'top-level')
jp2h_lst = [box for box in boxes if box.box_id == 'jp2h'] jp2h_lst = [box for box in boxes if box.box_id == 'jp2h']
jp2h = jp2h_lst[0] jp2h = jp2h_lst[0]
@ -1795,12 +1789,12 @@ def _validate_jp2h(boxes):
raise IOError(msg) raise IOError(msg)
colr = jp2h.box[colr_lst[0]] colr = jp2h.box[colr_lst[0]]
_validate_channel_definition(jp2h, colr) self._validate_channel_definition(jp2h, colr)
def _validate_channel_definition(self, jp2h, colr):
def _validate_channel_definition(jp2h, colr):
"""Validate the channel definition box.""" """Validate the channel definition box."""
cdef_lst = [j for (j, box) in enumerate(jp2h.box) if box.box_id == 'cdef'] cdef_lst = [j for (j, box) in enumerate(jp2h.box)
if box.box_id == 'cdef']
if len(cdef_lst) > 1: if len(cdef_lst) > 1:
msg = "Only one channel definition box is allowed in the " msg = "Only one channel definition box is allowed in the "
msg += "JP2 header." msg += "JP2 header."
@ -1820,12 +1814,10 @@ def _validate_channel_definition(jp2h, colr):
msg += "channel definition box." msg += "channel definition box."
raise IOError(msg) raise IOError(msg)
def _check_jp2h_child_boxes(self, boxes, parent_box_name):
JP2H_CHILDREN = set(['bpcc', 'cdef', 'cmap', 'ihdr', 'pclr'])
def _check_jp2h_child_boxes(boxes, parent_box_name):
"""Certain boxes can only reside in the JP2 header.""" """Certain boxes can only reside in the JP2 header."""
JP2H_CHILDREN = set(['bpcc', 'cdef', 'cmap', 'ihdr', 'pclr'])
box_ids = set([box.box_id for box in boxes]) box_ids = set([box.box_id for box in boxes])
intersection = box_ids.intersection(JP2H_CHILDREN) intersection = box_ids.intersection(JP2H_CHILDREN)
if len(intersection) > 0 and parent_box_name not in ['jp2h', 'jpch']: if len(intersection) > 0 and parent_box_name not in ['jp2h', 'jpch']:
@ -1835,27 +1827,24 @@ def _check_jp2h_child_boxes(boxes, parent_box_name):
# Recursively check any contained superboxes. # Recursively check any contained superboxes.
for box in boxes: for box in boxes:
if hasattr(box, 'box'): if hasattr(box, 'box'):
_check_jp2h_child_boxes(box.box, box.box_id) self._check_jp2h_child_boxes(box.box, box.box_id)
def _collect_box_count(self, boxes):
def _collect_box_count(boxes):
"""Count the occurences of each box type.""" """Count the occurences of each box type."""
count = Counter([box.box_id for box in boxes]) count = Counter([box.box_id for box in boxes])
# Add the counts in the superboxes. # Add the counts in the superboxes.
for box in boxes: for box in boxes:
if hasattr(box, 'box'): if hasattr(box, 'box'):
count.update(_collect_box_count(box.box)) count.update(self._collect_box_count(box.box))
return count return count
TOP_LEVEL_ONLY_BOXES = set(['dtbl']) def _check_superbox_for_top_levels(self, boxes):
def _check_superbox_for_top_levels(boxes):
"""Several boxes can only occur at the top level.""" """Several boxes can only occur at the top level."""
# We are only looking at the boxes contained in a superbox, so if any of # We are only looking at the boxes contained in a superbox, so if any
# the blacklisted boxes show up here, it's an error. # of the blacklisted boxes show up here, it's an error.
TOP_LEVEL_ONLY_BOXES = set(['dtbl'])
box_ids = set([box.box_id for box in boxes]) box_ids = set([box.box_id for box in boxes])
intersection = box_ids.intersection(TOP_LEVEL_ONLY_BOXES) intersection = box_ids.intersection(TOP_LEVEL_ONLY_BOXES)
if len(intersection) > 0: if len(intersection) > 0:
@ -1865,17 +1854,16 @@ def _check_superbox_for_top_levels(boxes):
# Recursively check any contained superboxes. # Recursively check any contained superboxes.
for box in boxes: for box in boxes:
if hasattr(box, 'box'): if hasattr(box, 'box'):
_check_superbox_for_top_levels(box.box) self._check_superbox_for_top_levels(box.box)
def _validate_top_level(self, boxes):
def _validate_top_level(boxes):
"""Several boxes can only occur at the top level.""" """Several boxes can only occur at the top level."""
# Add the counts in the superboxes. # Add the counts in the superboxes.
for box in boxes: for box in boxes:
if hasattr(box, 'box'): if hasattr(box, 'box'):
_check_superbox_for_top_levels(box.box) self._check_superbox_for_top_levels(box.box)
count = _collect_box_count(boxes) count = self._collect_box_count(boxes)
# Which boxes occur more than once? # Which boxes occur more than once?
multiples = [box_id for box_id, bcount in count.items() if bcount > 1] multiples = [box_id for box_id, bcount in count.items() if bcount > 1]
if 'dtbl' in multiples: if 'dtbl' in multiples:
@ -1883,26 +1871,23 @@ def _validate_top_level(boxes):
# If there is one data reference box, then there must also be one ftbl. # If there is one data reference box, then there must also be one ftbl.
if 'dtbl' in count and 'ftbl' not in count: if 'dtbl' in count and 'ftbl' not in count:
msg = 'The presence of a data reference box requires the presence of ' msg = 'The presence of a data reference box requires the presence '
msg += 'a fragment table box as well.' msg += 'of a fragment table box as well.'
raise IOError(msg) raise IOError(msg)
def _validate_singletons(self, boxes):
def _validate_singletons(boxes):
"""Several boxes can only occur once.""" """Several boxes can only occur once."""
count = _collect_box_count(boxes) count = self._collect_box_count(boxes)
# Which boxes occur more than once? # Which boxes occur more than once?
multiples = [box_id for box_id, bcount in count.items() if bcount > 1] multiples = [box_id for box_id, bcount in count.items() if bcount > 1]
if 'dtbl' in multiples: if 'dtbl' in multiples:
raise IOError('There can only be one dtbl box in a file.') raise IOError('There can only be one dtbl box in a file.')
JPX_IDS = ['asoc', 'nlst'] def _validate_jpx_brand(self, boxes, brand):
def _validate_jpx_brand(boxes, brand):
""" """
If there is a JPX box then the brand must be 'jpx '. If there is a JPX box then the brand must be 'jpx '.
""" """
JPX_IDS = ['asoc', 'nlst']
for box in boxes: for box in boxes:
if box.box_id in JPX_IDS: if box.box_id in JPX_IDS:
if brand != 'jpx ': if brand != 'jpx ':
@ -1911,13 +1896,14 @@ def _validate_jpx_brand(boxes, brand):
raise RuntimeError(msg) raise RuntimeError(msg)
if hasattr(box, 'box') != 0: if hasattr(box, 'box') != 0:
# Same set of checks on any child boxes. # Same set of checks on any child boxes.
_validate_jpx_brand(box.box, brand) self._validate_jpx_brand(box.box, brand)
def _validate_jpx_compatibility(self, boxes, compatibility_list):
def _validate_jpx_compatibility(boxes, compatibility_list):
""" """
If there is a JPX box then the compatibility list must also contain 'jpx '. If there is a JPX box then the compatibility list must also contain
'jpx '.
""" """
JPX_IDS = ['asoc', 'nlst']
jpx_cl = set(compatibility_list) jpx_cl = set(compatibility_list)
for box in boxes: for box in boxes:
if box.box_id in JPX_IDS: if box.box_id in JPX_IDS:
@ -1927,10 +1913,9 @@ def _validate_jpx_compatibility(boxes, compatibility_list):
raise RuntimeError(msg) raise RuntimeError(msg)
if hasattr(box, 'box') != 0: if hasattr(box, 'box') != 0:
# Same set of checks on any child boxes. # Same set of checks on any child boxes.
_validate_jpx_compatibility(box.box, compatibility_list) self._validate_jpx_compatibility(box.box, compatibility_list)
def _validate_label(self, boxes):
def _validate_label(boxes):
""" """
Label boxes can only be inside association, codestream headers, or Label boxes can only be inside association, codestream headers, or
compositing layer header boxes. compositing layer header boxes.
@ -1940,11 +1925,12 @@ def _validate_label(boxes):
if hasattr(box, 'box'): if hasattr(box, 'box'):
for boxi in box.box: for boxi in box.box:
if boxi.box_id == 'lbl ': if boxi.box_id == 'lbl ':
msg = "A label box cannot be nested inside a {0} box." msg = "A label box cannot be nested inside a "
msg += "{0} box."
msg = msg.format(box.box_id) msg = msg.format(box.box_id)
raise IOError(msg) raise IOError(msg)
# Same set of checks on any child boxes. # Same set of checks on any child boxes.
_validate_label(box.box) self._validate_label(box.box)
# Setup the default callback handlers. See the callback functions subsection # Setup the default callback handlers. See the callback functions subsection