Skip to content
Merged
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
110 changes: 110 additions & 0 deletions specfile/sanitizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -446,7 +446,117 @@ def build_lua_char_class(chars):
ordered = ordered.replace("-", "") + "-"
return f"[{ordered}]"

def _stage_to_lua_fragment(stage):
"""Convert a single pipeline stage to Lua code that transforms 'v'."""
cut = parse_cut(stage)
if cut:
mode, start, end, delim = cut
if mode == "bytes":
return None
elif mode == "field":
esc = lua_string_escape(lua_pattern_escape(delim))
return (
f"do local t={{}} "
f'for f in v:gmatch("[^{esc}]+") do t[#t+1]=f end '
f'v=t[{start}] or "" end'
)
elif mode == "range":
esc = lua_string_escape(lua_pattern_escape(delim))
return (
f"do local t={{}} "
f'for f in v:gmatch("[^{esc}]+") do t[#t+1]=f end '
f'v=table.concat(t,"{lua_string_escape(delim)}",'
f"{start},math.min(#t,{end})) end"
)
Comment thread
nforro marked this conversation as resolved.
if _RE_TR_LOWER.match(stage):
return "v=v:lower()"
if _RE_TR_UPPER.match(stage):
return "v=v:upper()"
m = _RE_TR_DELETE.match(stage)
if m:
chars = m.group(1)
if len(chars) == 1:
pat = lua_pattern_escape(chars)
else:
pat = build_lua_char_class(chars)
return f'v=(v:gsub("{lua_string_escape(pat)}", ""))'
m = _RE_TR_DELETE_BARE.match(stage)
if m:
pat = lua_string_escape(lua_pattern_escape(m.group(1)))
return f'v=(v:gsub("{pat}", ""))'
m = _RE_TR_REPLACE.match(stage)
if m:
pat = lua_string_escape(lua_pattern_escape(m.group(1)))
repl = lua_string_escape(lua_gsub_repl_escape(m.group(2)))
return f'v=(v:gsub("{pat}", "{repl}"))'
m = _RE_AWK_F.match(stage)
if m:
delim = m.group(1)
print_args = m.group(2)
parts = _RE_AWK_FIELDS.findall(print_args)
if parts:
if not all(is_safe_for_expand(sep) for _, sep in parts if sep):
return None
lua_parts = []
for field_num, separator in parts:
if field_num:
lua_parts.append(f'(t[{field_num}] or "")')
elif separator is not None:
lua_parts.append(f'"{lua_string_escape(separator)}"')
Comment thread
nforro marked this conversation as resolved.
if lua_parts:
lua_expr = " .. ".join(lua_parts)
esc = lua_string_escape(lua_pattern_escape(delim))
return (
f"do local t={{}} "
f'for f in v:gmatch("[^{esc}]+") do t[#t+1]=f end '
f"v={lua_expr} end"
)
substs = parse_sed_substs(stage)
if substs:
if all(is_safe_for_expand(repl) for _, repl, _ in substs):
parts = []
for pattern, repl, is_global in substs:
esc_pat = lua_string_escape(sed_pattern_to_lua(pattern))
esc_repl = lua_string_escape(lua_gsub_repl_escape(repl))
count_arg = "" if is_global else ", 1"
parts.append(
f'v=(v:gsub("{esc_pat}", "{esc_repl}"{count_arg}))'
)
return " ".join(parts)
return None

def convert_string_op(expr, cmd):
# Handle pipelines: split cmd into stages and compose Lua
try:
tokens = shlex.split(cmd, posix=False)
except ValueError:
tokens = None
if tokens:
stages = []
current = []
for token in tokens:
if token == "|":
if current:
stages.append(" ".join(current))
current = []
else:
current.append(token)
if current:
stages.append(" ".join(current))
if len(stages) > 1:
esc_expr = lua_string_escape(expr)
fragments = []
for stage in stages:
fragment = _stage_to_lua_fragment(stage)
if fragment is None:
return None
fragments.append(fragment)
lua_code = f'local v=rpm.expand("{esc_expr}")'
for fragment in fragments:
lua_code += f" {fragment}"
lua_code += " print(v)"
return f"%{{lua:{lua_code}}}"

# -- cut --
cut = parse_cut(cmd)
if cut:
Expand Down
60 changes: 58 additions & 2 deletions tests/unit/test_sanitizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -290,6 +290,10 @@ def test_pipe_to_cut_bytes(body, expected):
"echo %{date} | tr -d -",
'%{lua:print((rpm.expand("%{date}"):gsub("%-", "")))}',
),
(
"echo %{version} | tr '|' '.'",
'%{lua:print((rpm.expand("%{version}"):gsub("|", ".")))}',
),
],
)
def test_pipe_to_tr(body, expected):
Expand Down Expand Up @@ -337,8 +341,10 @@ def test_chained_sed():
assert Sanitizer.sanitize_shell_expansion(
"echo %{tag} | sed -e 's|.00$||' | sed -e 's|\\.||g'"
) == (
'%{lua:local v=(rpm.expand("%{tag}"):gsub(".00$", "", 1))'
' print((v:gsub("%.", "")))}'
'%{lua:local v=rpm.expand("%{tag}")'
' v=(v:gsub(".00$", "", 1))'
' v=(v:gsub("%.", ""))'
" print(v)}"
)


Expand All @@ -351,6 +357,56 @@ def test_chained_sed_rpm_escaped():
)


@pytest.mark.parametrize(
"body, expected",
[
(
"echo %{version} | cut -d. -f1 | cut -d~ -f1",
'%{lua:local v=rpm.expand("%{version}")'
' do local t={} for f in v:gmatch("[^%.]+") do t[#t+1]=f end v=t[1] or "" end'
' do local t={} for f in v:gmatch("[^~]+") do t[#t+1]=f end v=t[1] or "" end'
" print(v)}",
),
(
"echo %{version} | cut -d. -f1 | tr A B",
'%{lua:local v=rpm.expand("%{version}")'
' do local t={} for f in v:gmatch("[^%.]+") do t[#t+1]=f end v=t[1] or "" end'
' v=(v:gsub("A", "B"))'
" print(v)}",
),
(
"echo %{version} | tr '~' '.' | cut -d. -f1",
'%{lua:local v=rpm.expand("%{version}")'
' v=(v:gsub("~", "."))'
' do local t={} for f in v:gmatch("[^%.]+") do t[#t+1]=f end v=t[1] or "" end'
" print(v)}",
),
(
"echo %{version} | cut -d. -f1 | unsupported_command",
"%{nil}",
),
],
)
def test_piped_commands(body, expected):
assert Sanitizer.sanitize_shell_expansion(body) == expected


@pytest.mark.parametrize(
"body, expected",
[
(
"c=%{version}; echo $c | cut -d. -f1 | cut -d~ -f1",
'%{lua:local v=rpm.expand("%{version}")'
' do local t={} for f in v:gmatch("[^%.]+") do t[#t+1]=f end v=t[1] or "" end'
' do local t={} for f in v:gmatch("[^~]+") do t[#t+1]=f end v=t[1] or "" end'
" print(v)}",
),
],
)
def test_var_piped_commands(body, expected):
assert Sanitizer.sanitize_shell_expansion(body) == expected


@pytest.mark.parametrize(
"body, expected",
[
Expand Down
Loading