Skip to content

Commit d89291d

Browse files
committed
allow finalized pipelines to override the seed
1 parent 5f4d510 commit d89291d

3 files changed

Lines changed: 27 additions & 12 deletions

File tree

refinery/lib/argformats.py

Lines changed: 19 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -596,25 +596,23 @@ class DelayedArgument(LazyEvaluation):
596596
_CMD_SPLIT_TOKEN = ':'
597597

598598
def __init__(self, expression: str, reverse: bool = False, seed=None):
599+
requires_seed = seed is None
599600
self.expression = expression
600601
self.modifiers = []
601602
self.finalized = False
602-
if seed is not None:
603-
if reverse:
604-
if not expression.startswith(':'):
605-
expression = F':{expression}'
606-
else:
607-
if not expression.endswith(':'):
608-
expression = F'{expression}:'
609603
while not self.finalized:
610-
name, arguments, newexpr = self._split_modifier(expression, reverse)
604+
name, arguments, newexpr = self._split_modifier(
605+
expression,
606+
reverse,
607+
requires_seed,
608+
)
611609
if not name or not self.handler.can_handle(name, *arguments):
612610
break
613611
self.modifiers.append((name, arguments))
614612
expression = newexpr
615613
if self.handler.terminates(name):
616614
self.finalized = True
617-
if seed is not None:
615+
if not requires_seed and not self.finalized:
618616
if expression:
619617
rt = 'reverse ' if reverse else ''
620618
raise ValueError(F'{rt}expression {self.expression} with seed {seed} was not fully parsed.')
@@ -654,15 +652,24 @@ def _split_expression(self, expression: str, reverse: bool = False) -> tuple[str
654652
return (head, tail) if reverse else (tail, head)
655653
return expression, None
656654

657-
def _split_modifier(self, expression: str, reverse: bool = False) -> tuple[str | None, tuple[str], str]:
655+
def _split_modifier(
656+
self,
657+
sequence: str,
658+
reverse: bool = False,
659+
requires_seed: bool = True,
660+
) -> tuple[str | None, tuple[str, ...], str]:
658661
brackets = 0
659662
name = None
660663
argoffset = 0
661664
arguments = ()
662665

663-
rest, expression = self._split_expression(expression, reverse)
666+
rest, expression = self._split_expression(sequence, reverse)
667+
664668
if expression is None:
665-
return name, arguments, rest
669+
if requires_seed:
670+
return name, arguments, rest
671+
expression, rest = rest, ''
672+
666673
name = expression
667674

668675
for k, character in enumerate(expression):

test/lib/test_argformats.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,6 +52,10 @@ def test_skip_first_character_of_cyclic_key(self):
5252
key = argformats.DelayedArgument('take[1:16]:cycle:KITTY')()
5353
self.assertEqual(key, B'ITTYKITTYKITTYK')
5454

55+
def test_with_seed(self):
56+
key = argformats.DelayedBinaryArgument('snip[:2]:snip[1:]', seed=B'FOO')()
57+
self.assertEqual(key, B'OO')
58+
5559
def test_itob(self):
5660
data = argformats.DelayedArgument('itob:take[:4]:accu[0x1337]:A')()
5761
self.assertEqual(data, bytes.fromhex('3713371337133713'))

test/units/meta/test_reduce.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,10 @@ def test_concatenation(self):
88
pl = L('emit 5 4 3 2 1 0 [| reduce cca[var:t] ]')
99
self.assertEqual(pl(), B'012345')
1010

11+
def test_finalized_pipeline(self):
12+
pl = L('emit 5 4 3 2 1 0 [| reduce t:var ]')
13+
self.assertEqual(pl(), B'5')
14+
1115
def test_variables_are_retained(self):
1216
pl = L('emit +Y X X X [| put q index | reduce pf[{q}{}{t}] ]')
1317
self.assertEqual(pl(), B'3X2X1X+Y')

0 commit comments

Comments
 (0)