From 36ad34a555cbb1f8ce8cab5decc47f057d9b29d3 Mon Sep 17 00:00:00 2001 From: bbedward Date: Sun, 5 Jul 2026 17:16:55 -0400 Subject: [PATCH] keybinds: fix lua parsing on Hyprland for non-static configuration fixes #2753 --- core/internal/keybinds/providers/hyprland.go | 2 +- .../keybinds/providers/hyprland_lua_expand.go | 414 ++++++++++++++++++ .../providers/hyprland_lua_expand_test.go | 134 ++++++ .../keybinds/providers/hyprland_parser.go | 2 +- 4 files changed, 550 insertions(+), 2 deletions(-) create mode 100644 core/internal/keybinds/providers/hyprland_lua_expand.go create mode 100644 core/internal/keybinds/providers/hyprland_lua_expand_test.go diff --git a/core/internal/keybinds/providers/hyprland.go b/core/internal/keybinds/providers/hyprland.go index e5cd6733d..1108a2b87 100644 --- a/core/internal/keybinds/providers/hyprland.go +++ b/core/internal/keybinds/providers/hyprland.go @@ -1154,7 +1154,7 @@ func readLuaOrHyprlangOverride(path string) (map[string]*hyprlandOverrideBind, e if err != nil { return nil, err } - lines := strings.Split(string(data), "\n") + lines := expandLuaConfigLines(strings.Split(string(data), "\n")) parser := NewHyprlandParser("") pendingUnbinds := make(map[string]string) for _, line := range lines { diff --git a/core/internal/keybinds/providers/hyprland_lua_expand.go b/core/internal/keybinds/providers/hyprland_lua_expand.go new file mode 100644 index 000000000..e4715ecbd --- /dev/null +++ b/core/internal/keybinds/providers/hyprland_lua_expand.go @@ -0,0 +1,414 @@ +package providers + +import ( + "maps" + "regexp" + "strconv" + "strings" +) + +// Lua configs can express binds dynamically: variables (mainMod .. " + C"), +// tostring() calls, and numeric for loops (workspace binds). This resolves such +// expressions to literal key combos so the static bind parser can read them. + +const luaMaxLoopIterations = 1000 + +var ( + luaAssignRE = regexp.MustCompile(`^(?:local\s+)?([A-Za-z_][A-Za-z0-9_]*)\s*=\s*(.+)$`) + luaForRE = regexp.MustCompile(`^for\s+([A-Za-z_][A-Za-z0-9_]*)\s*=\s*(-?\d+)\s*,\s*(-?\d+)\s*(?:,\s*(-?\d+)\s*)?do\b(.*)$`) + luaTostringRE = regexp.MustCompile(`to(?:string|number)\s*\(\s*("(?:\\.|[^"])*"|'(?:\\.|[^'])*'|-?\d+(?:\.\d+)?)\s*\)`) + luaBlockOpenRE = regexp.MustCompile(`\b(?:function|for|while|if)\b`) + luaBlockCloseRE = regexp.MustCompile(`\bend\b`) + luaNumberRE = regexp.MustCompile(`^-?\d+(?:\.\d+)?$`) +) + +type luaVarEnv map[string]string + +type luaForHeader struct { + varName string + start int + stop int + step int + inline string +} + +func expandLuaConfigLines(lines []string) []string { + return expandLuaBlock(lines, luaVarEnv{}) +} + +func expandLuaBlock(lines []string, env luaVarEnv) []string { + out := make([]string, 0, len(lines)) + for i := 0; i < len(lines); i++ { + code := strings.TrimSpace(luaStripLineComment(lines[i])) + + if name, value, ok := parseLuaStringAssignment(code, env); ok { + env[name] = value + out = append(out, lines[i]) + continue + } + + header, ok := parseLuaForHeader(code) + if !ok { + out = append(out, resolveLuaDynamicLine(lines[i], env)) + continue + } + + body, consumed, complete := collectLuaForBody(header, lines, i) + if !complete { + out = append(out, resolveLuaDynamicLine(lines[i], env)) + continue + } + out = append(out, expandLuaForLoop(header, body, env)...) + i = consumed + } + return out +} + +func parseLuaStringAssignment(code string, env luaVarEnv) (name, value string, ok bool) { + m := luaAssignRE.FindStringSubmatch(code) + if m == nil { + return "", "", false + } + value, ok = evalLuaConcat(m[2], env) + if !ok { + return "", "", false + } + return m[1], value, true +} + +func parseLuaForHeader(code string) (luaForHeader, bool) { + m := luaForRE.FindStringSubmatch(code) + if m == nil { + return luaForHeader{}, false + } + start, _ := strconv.Atoi(m[2]) + stop, _ := strconv.Atoi(m[3]) + step := 1 + if m[4] != "" { + step, _ = strconv.Atoi(m[4]) + } + if step == 0 { + return luaForHeader{}, false + } + return luaForHeader{varName: m[1], start: start, stop: stop, step: step, inline: strings.TrimSpace(m[5])}, true +} + +func collectLuaForBody(header luaForHeader, lines []string, headerIdx int) (body []string, consumed int, complete bool) { + depth := 1 + if header.inline != "" { + delta := luaBlockDelta(header.inline) + if depth+delta <= 0 { + if stmt := strings.TrimSpace(strings.TrimSuffix(strings.TrimSpace(header.inline), "end")); stmt != "" { + body = append(body, stmt) + } + return body, headerIdx, true + } + depth += delta + body = append(body, header.inline) + } + for j := headerIdx + 1; j < len(lines); j++ { + delta := luaBlockDelta(lines[j]) + if depth+delta <= 0 { + return body, j, true + } + depth += delta + body = append(body, lines[j]) + } + return nil, headerIdx, false +} + +func expandLuaForLoop(header luaForHeader, body []string, env luaVarEnv) []string { + var out []string + inRange := func(v int) bool { + if header.step > 0 { + return v <= header.stop + } + return v >= header.stop + } + count := 0 + for v := header.start; inRange(v); v += header.step { + if count++; count > luaMaxLoopIterations { + break + } + value := strconv.Itoa(v) + iterLines := make([]string, len(body)) + for k, bl := range body { + iterLines[k] = substituteLuaIdent(bl, header.varName, value) + } + out = append(out, expandLuaBlock(iterLines, cloneLuaEnv(env))...) + } + return out +} + +func luaBlockDelta(line string) int { + masked := luaMaskStrings(line) + return len(luaBlockOpenRE.FindAllString(masked, -1)) - len(luaBlockCloseRE.FindAllString(masked, -1)) +} + +func resolveLuaDynamicLine(line string, env luaVarEnv) string { + if !strings.Contains(line, "hl.bind") && !strings.Contains(line, "hl.unbind") { + return line + } + line = normalizeLuaToString(line) + return rewriteLuaBindKeyArg(line, env) +} + +func normalizeLuaToString(line string) string { + return luaTostringRE.ReplaceAllStringFunc(line, func(m string) string { + inner := luaTostringRE.FindStringSubmatch(m)[1] + if inner[0] == '"' || inner[0] == '\'' { + return inner + } + return strconv.Quote(inner) + }) +} + +func rewriteLuaBindKeyArg(line string, env luaVarEnv) string { + for _, fn := range []string{"hl.bind", "hl.unbind"} { + idx := strings.Index(line, fn) + if idx < 0 { + continue + } + open := skipLuaWS(line, idx+len(fn)) + if open >= len(line) || line[open] != '(' { + continue + } + argStart := skipLuaWS(line, open+1) + expr, end, ok := parseLuaFirstArgExpr(line, argStart) + if !ok || isLuaPlainStringArg(expr) { + continue + } + value, ok := evalLuaConcat(expr, env) + if !ok { + continue + } + return line[:argStart] + strconv.Quote(value) + line[end:] + } + return line +} + +func evalLuaConcat(expr string, env luaVarEnv) (string, bool) { + parts := splitLuaConcat(expr) + var sb strings.Builder + for _, part := range parts { + value, ok := evalLuaOperand(part, env) + if !ok { + return "", false + } + sb.WriteString(value) + } + return sb.String(), true +} + +func evalLuaOperand(op string, env luaVarEnv) (string, bool) { + op = strings.TrimSpace(op) + if op == "" { + return "", false + } + switch op[0] { + case '"', '\'': + s, next, ok := parseLuaStringLiteral(op, 0) + return s, next == len(op) && ok + } + if luaNumberRE.MatchString(op) { + return op, true + } + if inner, ok := luaUnwrapCall(op, "tostring"); ok { + return evalLuaConcat(inner, env) + } + if inner, ok := luaUnwrapCall(op, "tonumber"); ok { + return evalLuaConcat(inner, env) + } + if value, ok := env[op]; ok { + return value, true + } + return "", false +} + +func splitLuaConcat(expr string) []string { + var parts []string + parenDepth, braceDepth, bracketDepth := 0, 0, 0 + inStr := byte(0) + esc := false + start := 0 + for i := 0; i < len(expr); i++ { + c := expr[i] + if inStr != 0 { + switch { + case esc: + esc = false + case c == '\\' && inStr == '"': + esc = true + case c == inStr: + inStr = 0 + } + continue + } + switch c { + case '"', '\'': + inStr = c + case '(': + parenDepth++ + case ')': + if parenDepth > 0 { + parenDepth-- + } + case '{': + braceDepth++ + case '}': + if braceDepth > 0 { + braceDepth-- + } + case '[': + bracketDepth++ + case ']': + if bracketDepth > 0 { + bracketDepth-- + } + case '.': + if parenDepth == 0 && braceDepth == 0 && bracketDepth == 0 && i+1 < len(expr) && expr[i+1] == '.' { + parts = append(parts, expr[start:i]) + i++ + start = i + 1 + } + } + } + return append(parts, expr[start:]) +} + +func substituteLuaIdent(line, name, value string) string { + if !strings.Contains(line, name) { + return line + } + var sb strings.Builder + inStr := byte(0) + esc := false + for i := 0; i < len(line); { + c := line[i] + if inStr != 0 { + sb.WriteByte(c) + switch { + case esc: + esc = false + case c == '\\' && inStr == '"': + esc = true + case c == inStr: + inStr = 0 + } + i++ + continue + } + if c == '"' || c == '\'' { + inStr = c + sb.WriteByte(c) + i++ + continue + } + if isLuaIdentStart(c) { + j := i + 1 + for j < len(line) && isLuaIdentByte(line[j]) { + j++ + } + word := line[i:j] + if word == name && (i == 0 || line[i-1] != '.') { + sb.WriteString(value) + } else { + sb.WriteString(word) + } + i = j + continue + } + sb.WriteByte(c) + i++ + } + return sb.String() +} + +func luaMaskStrings(line string) string { + b := []byte(line) + inStr := byte(0) + esc := false + for i := 0; i < len(b); i++ { + c := b[i] + if inStr != 0 { + wasEnd := !esc && c == inStr + esc = !esc && c == '\\' && inStr == '"' + b[i] = ' ' + if wasEnd { + inStr = 0 + } + continue + } + switch c { + case '"', '\'': + inStr = c + b[i] = ' ' + case '-': + if i+1 < len(b) && b[i+1] == '-' { + for ; i < len(b); i++ { + b[i] = ' ' + } + } + } + } + return string(b) +} + +func luaStripLineComment(line string) string { + inStr := byte(0) + esc := false + for i := 0; i+1 < len(line); i++ { + c := line[i] + if inStr != 0 { + switch { + case esc: + esc = false + case c == '\\' && inStr == '"': + esc = true + case c == inStr: + inStr = 0 + } + continue + } + switch c { + case '"', '\'': + inStr = c + case '-': + if line[i+1] == '-' { + return line[:i] + } + } + } + return line +} + +func luaUnwrapCall(op, fn string) (string, bool) { + op = strings.TrimSpace(op) + if !strings.HasPrefix(op, fn) { + return "", false + } + rest := strings.TrimSpace(op[len(fn):]) + if !strings.HasPrefix(rest, "(") || !strings.HasSuffix(rest, ")") { + return "", false + } + return rest[1 : len(rest)-1], true +} + +func isLuaPlainStringArg(expr string) bool { + expr = strings.TrimSpace(expr) + if expr == "" || (expr[0] != '"' && expr[0] != '\'') { + return false + } + _, next, ok := parseLuaStringLiteral(expr, 0) + return ok && next == len(expr) +} + +func isLuaIdentStart(c byte) bool { + return c == '_' || (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') +} + +func cloneLuaEnv(env luaVarEnv) luaVarEnv { + clone := make(luaVarEnv, len(env)) + maps.Copy(clone, env) + return clone +} diff --git a/core/internal/keybinds/providers/hyprland_lua_expand_test.go b/core/internal/keybinds/providers/hyprland_lua_expand_test.go new file mode 100644 index 000000000..9361e0e22 --- /dev/null +++ b/core/internal/keybinds/providers/hyprland_lua_expand_test.go @@ -0,0 +1,134 @@ +package providers + +import ( + "strings" + "testing" +) + +func TestExpandLuaConfigLinesVariableConcat(t *testing.T) { + lines := []string{ + `local mainMod = "SUPER"`, + `hl.bind(mainMod .. " + C", hl.dsp.window.close())`, + `hl.bind(mainMod .. " + H", hl.dsp.focus({direction = "l"}))`, + `hl.bind("ALT + TAB", hl.dsp.window.cycle_next({}))`, + } + got := expandLuaConfigLines(lines) + + want := []string{ + `hl.bind("SUPER + C",`, + `hl.bind("SUPER + H",`, + `hl.bind("ALT + TAB",`, + } + joined := strings.Join(got, "\n") + for _, w := range want { + if !strings.Contains(joined, w) { + t.Errorf("expanded output missing %q\n---\n%s", w, joined) + } + } +} + +func TestExpandLuaConfigLinesForLoop(t *testing.T) { + lines := []string{ + `local mainMod = "SUPER"`, + `for i = 1, 3 do`, + ` hl.bind(mainMod .. " + " .. i, hl.dsp.focus({workspace = tostring(i)}))`, + ` hl.bind(mainMod .. " SHIFT + " .. i, hl.dsp.window.move({workspace = tostring(i)}))`, + `end`, + } + got := strings.Join(expandLuaConfigLines(lines), "\n") + + for _, w := range []string{ + `hl.bind("SUPER + 1",`, + `hl.bind("SUPER + 2",`, + `hl.bind("SUPER + 3",`, + `hl.bind("SUPER SHIFT + 3",`, + `{workspace = "1"}`, + } { + if !strings.Contains(got, w) { + t.Errorf("expanded loop missing %q\n---\n%s", w, got) + } + } +} + +func TestParseLuaLinesDynamicBinds(t *testing.T) { + content := strings.Join([]string{ + `local mainMod = "SUPER"`, + `hl.bind(mainMod .. " + C", hl.dsp.window.close())`, + `hl.bind("ALT + TAB", hl.dsp.window.cycle_next({}))`, + `for i = 1, 2 do`, + ` hl.bind(mainMod .. " + " .. i, hl.dsp.focus({workspace = tostring(i)}))`, + `end`, + }, "\n") + + parser := NewHyprlandParser("") + section, err := parser.parseLuaLines(content, "", "test.lua", "") + if err != nil { + t.Fatalf("parseLuaLines: %v", err) + } + + keys := map[string]*HyprlandKeyBinding{} + for i := range section.Keybinds { + kb := §ion.Keybinds[i] + keys[parser.formatBindKey(kb)] = kb + } + + for _, want := range []string{"SUPER+C", "ALT+TAB", "SUPER+1", "SUPER+2"} { + if _, ok := keys[want]; !ok { + t.Errorf("missing bind %q; got %v", want, keysList(keys)) + } + } + if kb := keys["SUPER+C"]; kb != nil && kb.Dispatcher != "killactive" { + t.Errorf("SUPER+C dispatcher = %q, want killactive", kb.Dispatcher) + } + if kb := keys["SUPER+1"]; kb != nil { + if kb.Dispatcher != "workspace" || kb.Params != "1" { + t.Errorf("SUPER+1 = %q %q, want workspace 1", kb.Dispatcher, kb.Params) + } + } +} + +func keysList(m map[string]*HyprlandKeyBinding) []string { + out := make([]string, 0, len(m)) + for k := range m { + out = append(out, k) + } + return out +} + +func TestEvalLuaConcat(t *testing.T) { + env := luaVarEnv{"mainMod": "SUPER", "i": "5"} + tests := []struct { + expr string + want string + ok bool + }{ + {`mainMod .. " + C"`, "SUPER + C", true}, + {`mainMod .. " + " .. i`, "SUPER + 5", true}, + {`mainMod .. " + " .. tostring(i)`, "SUPER + 5", true}, + {`"ALT + TAB"`, "ALT + TAB", true}, + {`mainMod .. someFunc()`, "", false}, + {`unknownVar .. "x"`, "", false}, + } + for _, tt := range tests { + got, ok := evalLuaConcat(tt.expr, env) + if ok != tt.ok || (ok && got != tt.want) { + t.Errorf("evalLuaConcat(%q) = %q,%v want %q,%v", tt.expr, got, ok, tt.want, tt.ok) + } + } +} + +func TestSubstituteLuaIdent(t *testing.T) { + tests := []struct { + line, name, value, want string + }{ + {`hl.bind(m .. " + " .. i, x)`, "i", "3", `hl.bind(m .. " + " .. 3, x)`}, + {`hl.dsp.exec("light -i")`, "i", "3", `hl.dsp.exec("light -i")`}, + {`foo.i`, "i", "3", `foo.i`}, + {`tostring(i)`, "i", "3", `tostring(3)`}, + } + for _, tt := range tests { + if got := substituteLuaIdent(tt.line, tt.name, tt.value); got != tt.want { + t.Errorf("substituteLuaIdent(%q,%q,%q) = %q, want %q", tt.line, tt.name, tt.value, got, tt.want) + } + } +} diff --git a/core/internal/keybinds/providers/hyprland_parser.go b/core/internal/keybinds/providers/hyprland_parser.go index 57d8efc57..097b8e0ff 100644 --- a/core/internal/keybinds/providers/hyprland_parser.go +++ b/core/internal/keybinds/providers/hyprland_parser.go @@ -623,7 +623,7 @@ func (p *HyprlandParser) parseLuaLines(content string, baseDir, absPath, section prevSource := p.currentSource p.currentSource = absPath - lines := strings.Split(content, "\n") + lines := expandLuaConfigLines(strings.Split(content, "\n")) boundInFile := make(map[string]bool) for _, line := range lines { trimmed := strings.TrimSpace(line)