import fnmatch
import re
import string
import sys
from ClusterShell.Defaults import config_paths, DEFAULTS
import ClusterShell.NodeUtils as NodeUtils
from ClusterShell.RangeSet import RangeSet, RangeSetND, AUTOSTEP_DISABLED
from ClusterShell.RangeSet import RangeSetParseError
try:
basestring
except NameError:
basestring = str
DEF_GROUPS_CONFIGS = config_paths('groups.conf')
ILLEGAL_GROUP_CHARS = set("@,!&^*")
_DEF_RESOLVER_STD_GROUP = NodeUtils.GroupResolverConfig(DEF_GROUPS_CONFIGS,
ILLEGAL_GROUP_CHARS)
RESOLVER_STD_GROUP = _DEF_RESOLVER_STD_GROUP
RESOLVER_NOGROUP = -1
RESOLVER_NOINIT = -2
STD_GROUP_RESOLVER = RESOLVER_STD_GROUP
NOGROUP_RESOLVER = RESOLVER_NOGROUP
class NodeSetException(Exception):
class NodeSetError(NodeSetException):
class NodeSetParseError(NodeSetError):
def __init__(self, part, msg):
if part:
msg = "%s: \"%s\"" % (msg, part)
NodeSetError.__init__(self, msg)
self.part = part
class NodeSetParseRangeError(NodeSetParseError):
def __init__(self, rset_exc):
NodeSetParseError.__init__(self, str(rset_exc), "bad range")
class NodeSetExternalError(NodeSetError):
class NodeSetBase(object):
def __init__(self, pattern=None, rangeset=None, copy_rangeset=True,
autostep=None, fold_axis=None):
self._autostep = autostep
self._length = 0
self._patterns = {}
self.fold_axis = fold_axis if self.fold_axis is None and DEFAULTS.fold_axis:
self.fold_axis = DEFAULTS.fold_axis if pattern:
self._add(pattern, rangeset, copy_rangeset)
elif rangeset:
raise ValueError("missing pattern")
def get_autostep(self):
return self._autostep
def set_autostep(self, val):
if val is None:
self._autostep = None
else:
self._autostep = min(int(val), AUTOSTEP_DISABLED)
for pat, rset in self._patterns.items():
if rset:
rset.autostep = self._autostep
autostep = property(get_autostep, set_autostep)
def _iter(self):
for pat, rset in sorted(self._patterns.items()):
if rset:
autostep = rset.autostep
if rset.dim() == 1:
assert isinstance(rset, RangeSet)
for idx in rset:
yield pat, (idx,), autostep
else:
for rvec in rset:
yield pat, rvec, autostep
else:
yield pat, None, None
def _iterbase(self):
for pat, ivec, autostep in self._iter():
rset = None if ivec is not None:
assert len(ivec) > 0
if len(ivec) == 1:
rset = RangeSet.fromone(ivec[0], autostep)
else:
rset = RangeSetND([ivec], autostep)
yield NodeSetBase(pat, rset)
def __iter__(self):
for pat, ivec, _ in self._iter():
if ivec is not None:
yield pat % ivec
else:
yield pat % ()
striter = __iter__
def nsiter(self):
for pat, ivec, autostep in self._iter():
nodeset = self.__class__()
if ivec is not None:
if len(ivec) == 1:
nodeset._add_new(pat, RangeSet.fromone(ivec[0]))
else:
nodeset._add_new(pat, RangeSetND([ivec], autostep))
else:
nodeset._add_new(pat, None)
yield nodeset
def contiguous(self):
for pat, rangeset in sorted(self._patterns.items()):
if rangeset:
for cont_rset in rangeset.contiguous():
nodeset = self.__class__()
nodeset._add_new(pat, cont_rset)
yield nodeset
else:
nodeset = self.__class__()
nodeset._add_new(pat, None)
yield nodeset
def __len__(self):
cnt = 0
for rangeset in self._patterns.values():
if rangeset:
cnt += len(rangeset)
else:
cnt += 1
return cnt
def _iter_nd_pat(self, pat, rset):
try:
dimcnt = rset.dim()
if self.fold_axis is None:
fold_axis = range(dimcnt)
else:
fold_axis = [int(x) % dimcnt for x in self.fold_axis
if -dimcnt <= int(x) < dimcnt]
except (TypeError, ValueError) as exc:
raise NodeSetParseError("fold_axis=%s" % self.fold_axis, exc)
for rgvec in rset.vectors():
rgnargs = [] for axis, rangeset in enumerate(rgvec):
if len(rangeset) > 1:
if axis not in fold_axis: rgstrit = rangeset.striter()
else:
rgstrit = ["[%s]" % rangeset]
else:
rgstrit = [str(rangeset)]
t_rgnargs = []
for rgstr in rgstrit: if not rgnargs:
t_rgnargs.append([rgstr])
else:
for rga in rgnargs:
t_rgnargs.append(rga + [rgstr])
rgnargs = t_rgnargs
for rgargs in rgnargs:
yield pat % tuple(rgargs)
def __str__(self):
results = []
try:
for pat, rset in sorted(self._patterns.items()):
if not rset:
results.append(pat % ())
elif rset.dim() == 1:
if self.fold_axis is None or \
list(x for x in self.fold_axis if -1 <= int(x) < 1):
rgs = str(rset)
cnt = len(rset)
if cnt > 1:
rgs = "[%s]" % rgs
results.append(pat % rgs)
else:
results.extend((pat % rgs for rgs in rset.striter()))
elif rset.dim() > 1:
results.extend(self._iter_nd_pat(pat, rset))
except TypeError:
raise NodeSetParseError(pat, "Internal error: node pattern and "
"ranges mismatch")
return ",".join(results)
def copy(self):
cpy = self.__class__()
cpy.fold_axis = self.fold_axis
cpy._autostep = self._autostep
cpy._length = self._length
dic = {}
for pat, rangeset in self._patterns.items():
if rangeset is None:
dic[pat] = None
else:
dic[pat] = rangeset.copy()
cpy._patterns = dic
return cpy
def __contains__(self, other):
return self.issuperset(other)
def _binary_sanity_check(self, other):
if not isinstance(other, NodeSetBase):
raise TypeError("Binary operation only permitted between "
"NodeSetBase")
def issubset(self, other):
self._binary_sanity_check(other)
return other.issuperset(self)
def issuperset(self, other):
self._binary_sanity_check(other)
status = True
for pat, erangeset in other._patterns.items():
rangeset = self._patterns.get(pat)
if rangeset:
status = rangeset.issuperset(erangeset)
else:
status = pat in self._patterns
if not status:
break
return status
def __eq__(self, other):
if not isinstance(other, NodeSetBase):
return NotImplemented
return len(self) == len(other) and self.issuperset(other)
__le__ = issubset
__ge__ = issuperset
def __lt__(self, other):
self._binary_sanity_check(other)
return len(self) < len(other) and self.issubset(other)
def __gt__(self, other):
self._binary_sanity_check(other)
return len(self) > len(other) and self.issuperset(other)
def _extractslice(self, index):
length = len(self)
if index.start is None:
sl_start = 0
elif index.start < 0:
sl_start = max(0, length + index.start)
else:
sl_start = index.start
if index.stop is None:
sl_stop = sys.maxsize
elif index.stop < 0:
sl_stop = max(0, length + index.stop)
else:
sl_stop = index.stop
if index.step is None:
sl_step = 1
elif index.step < 0:
if index.start is not None or index.stop is not None:
raise IndexError("illegal start and stop when negative step "
"is used")
stepmod = (length + -index.step - 1) % -index.step
if stepmod > 0:
sl_start += stepmod
sl_step = -index.step
else:
sl_step = index.step
if not isinstance(sl_start, int) or not isinstance(sl_stop, int) \
or not isinstance(sl_step, int):
raise TypeError("slice indices must be integers")
return sl_start, sl_stop, sl_step
def __getitem__(self, index):
if isinstance(index, slice):
inst = NodeSetBase()
sl_start, sl_stop, sl_step = self._extractslice(index)
sl_next = sl_start
if sl_stop <= sl_next:
return inst
length = 0
for pat, rangeset in sorted(self._patterns.items()):
if rangeset:
cnt = len(rangeset)
offset = sl_next - length
if offset < cnt:
num = min(sl_stop - sl_next, cnt - offset)
inst._add(pat, rangeset[offset:offset + num:sl_step])
else:
length += cnt
continue
else:
cnt = num = 1
if sl_next > length:
length += cnt
continue
inst._add(pat, None)
sl_next += num
if (sl_next - sl_start) % sl_step:
sl_next = sl_start + \
((sl_next - sl_start)//sl_step + 1) * sl_step
if sl_next >= sl_stop:
break
length += cnt
return inst
elif isinstance(index, int):
if index < 0:
length = len(self)
if index >= -length:
index = length + index else:
raise IndexError("%d out of range" % index)
length = 0
for pat, rangeset in sorted(self._patterns.items()):
if rangeset:
cnt = len(rangeset)
if index < length + cnt:
if rangeset.dim() == 1:
return pat % rangeset[index-length:index-length+1]
else:
sub = rangeset[index-length:index-length+1]
for rgvec in sub.vectors():
return pat % (tuple(rgvec))
else:
cnt = 1
if index == length:
return pat
length += cnt
raise IndexError("%d out of range" % index)
else:
raise TypeError("NodeSet indices must be integers")
def _rangeset_index(self, rangeset, orangeset):
if rangeset is None or orangeset is None:
return 0 if rangeset is None and orangeset is None else None
try:
if isinstance(orangeset, RangeSetND):
return rangeset.index(next(iter(orangeset)))
return rangeset.index(next(orangeset.striter()))
except ValueError:
return None
def index(self, other, start=0, stop=None):
self._binary_sanity_check(other)
if len(other) != 1:
raise ValueError("index() argument must be a single node")
opat, orangeset = list(other._patterns.items())[0]
found = None
base = 0
for pat, rangeset in sorted(self._patterns.items()):
if pat == opat:
offset = self._rangeset_index(rangeset, orangeset)
if offset is not None:
found = base + offset
break
if rangeset:
base += len(rangeset)
else:
base += 1
if found is None:
raise ValueError("'%s' is not in nodeset" % other)
if start != 0 or stop is not None:
length = len(self)
if start < 0:
start = max(0, length + start)
if stop is None:
stop = length
elif stop < 0:
stop = max(0, length + stop)
if not start <= found < stop:
raise ValueError("'%s' is not in nodeset" % other)
return found
def _add_new(self, pat, rangeset):
assert pat not in self._patterns
self._patterns[pat] = rangeset
def _add(self, pat, rangeset, copy_rangeset=True):
if pat in self._patterns:
pat_e = self._patterns[pat]
if (pat_e is None) is not (rangeset is None):
raise NodeSetError("Invalid operation")
if pat_e:
pat_e.update(rangeset)
else:
if rangeset and copy_rangeset:
rangeset = rangeset.copy()
if self._autostep is not None:
rangeset.autostep = self._autostep
self._add_new(pat, rangeset)
def union(self, other):
self_copy = self.copy()
self_copy.update(other)
return self_copy
def __or__(self, other):
if not isinstance(other, NodeSetBase):
return NotImplemented
return self.union(other)
def add(self, other):
self.update(other)
def update(self, other):
for pat, rangeset in other._patterns.items():
self._add(pat, rangeset)
def updaten(self, others):
for other in others:
self.update(other)
def clear(self):
self._patterns.clear()
def __ior__(self, other):
self._binary_sanity_check(other)
self.update(other)
return self
def intersection(self, other):
self_copy = self.copy()
self_copy.intersection_update(other)
return self_copy
def __and__(self, other):
if not isinstance(other, NodeSet):
return NotImplemented
return self.intersection(other)
def intersection_update(self, other):
if other is self:
return
tmp_ns = NodeSetBase()
for pat, irangeset in other._patterns.items():
rangeset = self._patterns.get(pat)
if rangeset:
irset = rangeset.intersection(irangeset)
if len(irset) > 0:
tmp_ns._add(pat, irset, copy_rangeset=False)
elif not irangeset and pat in self._patterns:
tmp_ns._add(pat, None)
self._patterns = tmp_ns._patterns
def __iand__(self, other):
self._binary_sanity_check(other)
self.intersection_update(other)
return self
def difference(self, other):
self_copy = self.copy()
self_copy.difference_update(other)
return self_copy
def __sub__(self, other):
if not isinstance(other, NodeSetBase):
return NotImplemented
return self.difference(other)
def difference_update(self, other, strict=False):
purge_patterns = []
for pat, erangeset in other._patterns.items():
rangeset = self._patterns.get(pat)
if rangeset:
rangeset.difference_update(erangeset, strict)
if len(rangeset) == 0:
purge_patterns.append(pat)
else:
if pat in self._patterns:
purge_patterns.append(pat)
elif strict:
raise KeyError(pat)
for pat in purge_patterns:
del self._patterns[pat]
def __isub__(self, other):
self._binary_sanity_check(other)
self.difference_update(other)
return self
def remove(self, elem):
self.difference_update(elem, True)
def symmetric_difference(self, other):
self_copy = self.copy()
self_copy.symmetric_difference_update(other)
return self_copy
def __xor__(self, other):
if not isinstance(other, NodeSet):
return NotImplemented
return self.symmetric_difference(other)
def symmetric_difference_update(self, other):
purge_patterns = []
for pat, rangeset in self._patterns.items():
brangeset = other._patterns.get(pat)
if brangeset:
rangeset.symmetric_difference_update(brangeset)
else:
if pat in other._patterns:
purge_patterns.append(pat)
for pat, brangeset in other._patterns.items():
rangeset = self._patterns.get(pat)
if not rangeset and not pat in self._patterns:
self._add(pat, brangeset)
for pat, rangeset in self._patterns.items():
if rangeset is not None and len(rangeset) == 0:
purge_patterns.append(pat)
for pat in purge_patterns:
del self._patterns[pat]
def __ixor__(self, other):
self._binary_sanity_check(other)
self.symmetric_difference_update(other)
return self
def _strip_escape(nsstr):
return nsstr.strip().replace('%', '%%')
def _rsets4nsb(rsets, autostep):
if len(rsets) > 1:
return RangeSetND([rsets], None, autostep, copy_rangeset=False)
elif len(rsets) == 1:
return rsets[0]
class ParsingEngine(object):
OP_CODES = {',': 'update',
'!': 'difference_update',
'&': 'intersection_update',
'^': 'symmetric_difference_update'}
OP_CODES_PAT = '[%s]' % re.escape(''.join(OP_CODES.keys()))
BRACKET_OPEN = '['
BRACKET_CLOSE = ']'
def __init__(self, group_resolver, node_wildcard_enable=True):
self.group_resolver = group_resolver
self.base_node_re = re.compile(r"(\D*)(\d*)")
self.node_wc = node_wildcard_enable
def parse(self, nsobj, autostep):
if nsobj is None:
return NodeSetBase()
if isinstance(nsobj, NodeSetBase):
return nsobj
if isinstance(nsobj, basestring):
try:
return self.parse_string(str(nsobj), autostep)
except (NodeUtils.GroupSourceQueryFailed, RuntimeError) as exc:
raise NodeSetParseError(nsobj, str(exc))
raise TypeError("Unsupported NodeSet input %s" % type(nsobj))
def parse_string(self, nsstr, autostep, namespace=None):
alln_cache = None nodeset = NodeSetBase()
nsstr = _strip_escape(nsstr)
for opc, pat, rgnd in self._scan_string(nsstr, autostep):
if self.group_resolver and pat[0] == '@':
ns_group = NodeSetBase()
for nodegroup in NodeSetBase(pat, rgnd):
ns_str_ext, ns_nsp_ext = self.parse_group_string(nodegroup,
namespace)
if ns_str_ext: ns_group.update(self.parse_string(ns_str_ext,
autostep,
ns_nsp_ext))
getattr(nodeset, opc)(ns_group)
elif self.group_resolver and self.node_wc and ('*' in pat or
'?' in pat):
wcmasks = (str(wcn) for wcn in NodeSetBase(pat, rgnd, False))
if alln_cache is None:
self.node_wc = False try:
nsb = NodeSetBase()
for res in self.all_nodes(namespace):
nsb.update(self.parse_string(res, autostep,
namespace))
alln_cache = set(str(node) for node in nsb)
finally:
self.node_wc = True
alln = alln_cache.copy()
wcns = NodeSetBase()
for wcmask in wcmasks:
for node in fnmatch.filter(alln, wcmask):
alln.remove(node) wcp, wcr = self._scan_string_single(node, autostep)
wcrgnd = _rsets4nsb(wcr, autostep)
wcns.update(NodeSetBase(wcp, wcrgnd, False))
getattr(nodeset, opc)(wcns)
else:
getattr(nodeset, opc)(NodeSetBase(pat, rgnd, False))
return nodeset
def parse_string_single(self, nsstr, autostep):
pat, rangesets = self._scan_string_single(_strip_escape(nsstr),
autostep)
if len(rangesets) > 1:
rgobj = RangeSetND([rangesets], None, autostep, copy_rangeset=False)
elif len(rangesets) == 1:
rgobj = rangesets[0]
else: rgobj = None
return NodeSetBase(pat, rgobj, False)
def parse_group(self, group, namespace=None, autostep=None):
assert self.group_resolver is not None
nodestr = self.group_resolver.group_nodes(group, namespace)
return self.parse(",".join(nodestr), autostep)
def parse_group_string(self, nodegroup, namespace=None):
assert nodegroup[0] == '@'
assert self.group_resolver is not None
grpstr = group = nodegroup[1:]
if grpstr.find(':') >= 0:
namespace, group = grpstr.split(':', 1)
if group == '*': reslist = self.all_nodes(namespace)
elif group.startswith('@'): reslist = self.grouplist(grpstr[1:])
else:
reslist = self.group_resolver.group_nodes(group, namespace)
return ','.join(reslist), namespace
def grouplist(self, namespace=None):
grpset = NodeSetBase()
for grpstr in self.group_resolver.grouplist(namespace):
grpstr = _strip_escape(grpstr)
for opc, pat, rgnd in self._scan_string(grpstr, None):
getattr(grpset, opc)(NodeSetBase(pat, rgnd, False))
return list(grpset)
def all_nodes(self, namespace=None):
assert self.group_resolver is not None
alln = []
try:
alln = self.group_resolver.all_nodes(namespace)
except NodeUtils.GroupSourceNoUpcall:
try:
for grp in self.grouplist(namespace):
alln += self.group_resolver.group_nodes(grp, namespace)
except NodeUtils.GroupSourceNoUpcall:
msg = "Not enough working methods (all or map + list) to " \
"get all nodes"
raise NodeSetExternalError(msg)
except NodeUtils.GroupSourceQueryFailed as exc:
raise NodeSetExternalError("Failed to get all nodes: %s" % exc)
return alln
def _next_op(self, pat):
mobj = re.search(ParsingEngine.OP_CODES_PAT, pat)
if mobj:
return mobj.span()[0], mobj.group()
else:
return -1, None
def _scan_string_single(self, nsstr, autostep):
pfx_nd = [mobj.groups() for mobj in self.base_node_re.finditer(nsstr)]
pfx_nd = pfx_nd[:-1]
if not pfx_nd:
raise NodeSetParseError(nsstr, "parse error")
pat = ""
rangesets = []
for pfx, idx in pfx_nd:
if idx:
pad = 0
if int(idx) != 0:
idxs = idx.lstrip("0")
if len(idx) - len(idxs) > 0:
pad = len(idx)
idxint = int(idxs)
else:
if len(idx) > 1:
pad = len(idx)
idxint = 0
if idxint > 1e100:
raise NodeSetParseRangeError( \
RangeSetParseError(idx, "invalid rangeset index"))
pat += "%s%%s" % pfx
rangesets.append(RangeSet.fromone(idxint, pad, autostep))
else:
pat += pfx
return pat, rangesets
def _scan_string(self, nsstr, autostep):
next_op_code = ',' while nsstr:
nsstr = nsstr.lstrip()
rsets = []
op_code = next_op_code
op_idx, next_op_code = self._next_op(nsstr)
bracket_idx = nsstr.find(self.BRACKET_OPEN)
if bracket_idx >= 0 and (op_idx > bracket_idx or op_idx < 0):
newpat = ""
sfx = nsstr
while bracket_idx >= 0 and (op_idx > bracket_idx or op_idx < 0):
pfx, sfx = sfx.split(self.BRACKET_OPEN, 1)
try:
rng, sfx = sfx.split(self.BRACKET_CLOSE, 1)
except ValueError:
raise NodeSetParseError(nsstr, "missing bracket")
if pfx.find(self.BRACKET_CLOSE) > -1:
raise NodeSetParseError(pfx, "illegal closing bracket")
if len(sfx) > 0:
bra_end = sfx.find(self.BRACKET_CLOSE)
bra_start = sfx.find(self.BRACKET_OPEN)
if bra_start == -1:
bra_start = bra_end + 1
if bra_end >= 0 and bra_end < bra_start:
msg = "illegal closing bracket"
raise NodeSetParseError(sfx, msg)
pfxlen, sfxlen = len(pfx), len(sfx)
if sfxlen > 0:
try:
sfx, rng = self._amend_trailing_digits(sfx, rng)
except RangeSetParseError as ex:
raise NodeSetParseRangeError(ex)
if pfxlen > 0:
try:
pfx, rng = self._amend_leading_digits(pfx, rng)
except RangeSetParseError as ex:
raise NodeSetParseRangeError(ex)
if pfx:
pfx, pfxrvec = self._scan_string_single(pfx, autostep)
rsets += pfxrvec
bracket_idx = sfx.find(self.BRACKET_OPEN,
bracket_idx - pfxlen)
op_idx, next_op_code = self._next_op(sfx)
if len(sfx) > 0 and sfx[0] == '[':
msg = "illegal reopening bracket"
raise NodeSetParseError(sfx, msg)
newpat += "%s%%s" % pfx
try:
rsets.append(RangeSet(rng, autostep))
except RangeSetParseError as ex:
raise NodeSetParseRangeError(ex)
op_idx, next_op_code = self._next_op(sfx)
if op_idx < 0:
nsstr = None
else:
sfx, nsstr = sfx.split(next_op_code, 1)
if not nsstr:
msg = "missing nodeset operand with '%s' " \
"operator" % next_op_code
raise NodeSetParseError(None, msg)
sfx = sfx.rstrip()
if sfx:
sfx, sfxrvec = self._scan_string_single(sfx, autostep)
newpat += sfx
rsets += sfxrvec
else:
if op_idx < 0:
node = nsstr
nsstr = None else:
node, nsstr = nsstr.split(next_op_code, 1)
if not node or not nsstr:
msg = "missing nodeset operand with '%s' " \
"operator" % next_op_code
raise NodeSetParseError(node or nsstr, msg)
if node.find(self.BRACKET_CLOSE) > -1:
raise NodeSetParseError(node, "illegal closing bracket")
node = node.rstrip()
newpat, rsets = self._scan_string_single(node, autostep)
op = ParsingEngine.OP_CODES[op_code]
yield op, newpat, _rsets4nsb(rsets, autostep)
def _amend_leading_digits(self, outer, inner):
outerstrip = outer.rstrip(string.digits)
outerlen, outerstriplen = len(outer), len(outerstrip)
if outerstriplen < outerlen:
outerdigits = outer[outerstriplen:]
inner = ','.join(
'-'.join(outerdigits + bound for bound in elem.split('-'))
for elem in (str(subrng)
for subrng in RangeSet(inner).contiguous()))
return outerstrip, inner
def _amend_trailing_digits(self, outer, inner):
outerstrip = outer.lstrip(string.digits)
outerlen, outerstriplen = len(outer), len(outerstrip)
if outerstriplen < outerlen:
if '/' in inner:
msg = "illegal trailing digits after range with steps"
raise NodeSetParseError(outer, msg)
outerdigits = outer[0:outerlen-outerstriplen]
outlen = len(outerdigits)
def shiftstep(orig, power):
if '-' in orig:
return orig + '/1' + '0' * power
return orig inner = ','.join(shiftstep(s, outlen) for s in
('-'.join(bound + outerdigits
for bound in elem.split('-'))
for elem in inner.split(',')))
return outerstrip, inner
class NodeSet(NodeSetBase):
_VERSION = 2
def __init__(self, nodes=None, autostep=None, resolver=None,
fold_axis=None):
NodeSetBase.__init__(self, autostep=autostep, fold_axis=fold_axis)
if resolver in (RESOLVER_NOGROUP, RESOLVER_NOINIT):
self._resolver = None
else:
self._resolver = resolver or RESOLVER_STD_GROUP
if resolver == RESOLVER_NOINIT:
self._parser = None
else:
self._parser = ParsingEngine(self._resolver)
self.update(nodes)
@classmethod
def _fromlist1(cls, nodelist, autostep=None, resolver=None):
inst = NodeSet(autostep=autostep, resolver=resolver)
for single in nodelist:
inst.update(inst._parser.parse_string_single(single, autostep))
return inst
@classmethod
def fromlist(cls, nodelist, autostep=None, resolver=None):
inst = NodeSet(autostep=autostep, resolver=resolver)
inst.updaten(nodelist)
return inst
@classmethod
def fromall(cls, groupsource=None, autostep=None, resolver=None):
inst = NodeSet(autostep=autostep, resolver=resolver)
try:
if not inst._resolver:
raise NodeSetExternalError("Group resolver is not defined")
else:
inst.updaten(inst._parser.all_nodes(groupsource))
except NodeUtils.GroupResolverError as exc:
errmsg = "Group source error (%s: %s)" % (exc.__class__.__name__,
exc)
raise NodeSetExternalError(errmsg)
return inst
def __getstate__(self):
odict = self.__dict__.copy()
odict['_version'] = NodeSet._VERSION
del odict['_resolver']
del odict['_parser']
return odict
def __setstate__(self, dic):
self.__dict__.update(dic)
self._resolver = None
self._parser = ParsingEngine(None)
if getattr(self, '_version', 1) <= 1:
self.fold_axis = None
old_patterns = self._patterns
self._patterns = {}
for pat, rangeset in sorted(old_patterns.items()):
if rangeset:
assert isinstance(rangeset, RangeSet)
rgs = str(rangeset)
if len(rangeset) > 1:
rgs = "[%s]" % rgs
self.update(pat % rgs)
else:
self.update(pat)
def copy(self):
cpy = self.__class__(resolver=RESOLVER_NOINIT)
dic = {}
for pat, rangeset in self._patterns.items():
if rangeset is None:
dic[pat] = None
else:
dic[pat] = rangeset.copy()
cpy._patterns = dic
cpy.fold_axis = self.fold_axis
cpy._autostep = self._autostep
cpy._resolver = self._resolver
cpy._parser = self._parser
return cpy
__copy__ = copy
def _find_groups(self, node, namespace, allgroups):
if allgroups:
for grp, nodeset in allgroups.items():
if node in nodeset:
yield grp
else:
try:
for group in self._resolver.node_groups(node, namespace):
yield group
except NodeUtils.GroupSourceQueryFailed as exc:
msg = "Group source query failed: %s" % exc
raise NodeSetExternalError(msg)
def _groups2(self, groupsource=None, autostep=None):
if not self._resolver:
raise NodeSetExternalError("No node group resolver")
try:
allgrplist = self._parser.grouplist(groupsource)
except NodeUtils.GroupSourceError:
allgrplist = None
groups_info = {}
allgroups = {}
if self._resolver.has_node_groups(groupsource) and \
(not allgrplist or len(allgrplist) >= len(self)):
pass
else:
if not allgrplist: return groups_info try:
for grp in allgrplist:
nodelist = self._resolver.group_nodes(grp, groupsource)
allgroups[grp] = NodeSet(",".join(nodelist),
resolver=self._resolver)
except NodeUtils.GroupSourceQueryFailed as exc:
raise NodeSetExternalError("Unable to map a group " \
"previously listed\n\tFailed command: %s" % exc)
for node in self._iterbase():
for grp in self._find_groups(node, groupsource, allgroups):
if grp not in groups_info:
nodes = self._parser.parse_group(grp, groupsource, autostep)
groups_info[grp] = (1, nodes)
else:
i, nodes = groups_info[grp]
groups_info[grp] = (i + 1, nodes)
return groups_info
def groups(self, groupsource=None, noprefix=False):
groups = self._groups2(groupsource, self._autostep)
result = {}
for grp, (_, nsb) in groups.items():
if groupsource and not noprefix:
key = "@%s:%s" % (groupsource, grp)
else:
key = "@" + grp
result[key] = (NodeSet(nsb, resolver=self._resolver),
self.intersection(nsb))
return result
def regroup(self, groupsource=None, autostep=None, overlap=False,
noprefix=False):
groups = self._groups2(groupsource, autostep)
if not groups:
return str(self)
fulls = []
for k, (i, nodes) in groups.items():
assert i <= len(nodes)
if i == len(nodes):
fulls.append((i, k))
rest = NodeSet(self, resolver=RESOLVER_NOGROUP)
regrouped = NodeSet(resolver=RESOLVER_NOGROUP)
for _, grp in sorted(fulls, key=lambda x: (-x[0], x[1])):
if not overlap and groups[grp][1] not in rest:
continue
if groupsource and not noprefix:
regrouped.update("@%s:%s" % (groupsource, grp))
else:
regrouped.update("@" + grp)
rest.difference_update(groups[grp][1])
if not rest:
return str(regrouped)
if regrouped:
return "%s,%s" % (regrouped, rest)
return str(rest)
def issubset(self, other):
nodeset = self._parser.parse(other, self._autostep)
return NodeSetBase.issuperset(nodeset, self)
def issuperset(self, other):
nodeset = self._parser.parse(other, self._autostep)
return NodeSetBase.issuperset(self, nodeset)
def __getitem__(self, index):
base = NodeSetBase.__getitem__(self, index)
if not isinstance(base, NodeSetBase):
return base
inst = NodeSet(autostep=self._autostep, resolver=self._resolver)
inst._patterns = base._patterns
return inst
def index(self, other, start=0, stop=None):
nodeset = self._parser.parse(other, self._autostep)
return NodeSetBase.index(self, nodeset, start, stop)
def split(self, nbr):
assert(nbr > 0)
slice_size = len(self) // nbr
left = len(self) % nbr
begin = 0
for i in range(0, min(nbr, len(self))):
length = slice_size + int(i < left)
yield self[begin:begin + length]
begin += length
def update(self, other):
nodeset = self._parser.parse(other, self._autostep)
NodeSetBase.update(self, nodeset)
def intersection_update(self, other):
nodeset = self._parser.parse(other, self._autostep)
NodeSetBase.intersection_update(self, nodeset)
def difference_update(self, other, strict=False):
nodeset = self._parser.parse(other, self._autostep)
NodeSetBase.difference_update(self, nodeset, strict)
def symmetric_difference_update(self, other):
nodeset = self._parser.parse(other, self._autostep)
NodeSetBase.symmetric_difference_update(self, nodeset)
def expand(pat):
return list(NodeSet(pat))
def fold(pat):
return str(NodeSet(pat))
def grouplist(namespace=None, resolver=None):
return ParsingEngine(resolver or RESOLVER_STD_GROUP).grouplist(namespace)
def std_group_resolver():
return RESOLVER_STD_GROUP
def set_std_group_resolver(new_resolver):
global RESOLVER_STD_GROUP
RESOLVER_STD_GROUP = new_resolver or _DEF_RESOLVER_STD_GROUP
def set_std_group_resolver_config(groupsconf, illegal_chars=None):
if groupsconf:
if illegal_chars is None:
illegal_chars = ILLEGAL_GROUP_CHARS
group_resolver = NodeUtils.GroupResolverConfig(groupsconf,
illegal_chars)
set_std_group_resolver(group_resolver)