From 421b65bf43bd318cdd3801edd01ee1ba075ee099 Mon Sep 17 00:00:00 2001 From: Jo Date: Mon, 15 Jun 2026 15:25:40 +0200 Subject: [PATCH] feat: Add UpgradeSelect hook, logging and security example (#2855) * add upgrade hook structure * remove extra field * feat: Add UpgradeSelect hook logging and security example Add yay.log API to UpgradeSelect hooks for informational and warning messages. - Add yay.log.info() and yay.log.warn() for hook logging - Add doc/examples/recently_modified.lua: practical example to pre-exclude recently-modified AUR packages (supply chain attack mitigation) - Update doc/init.lua with improved UpgradeSelect example - Update doc/lua.md with logging API documentation - Refactor autocmd and load modules for cleaner log integration - Add comprehensive unit tests for logging functionality * lint --- doc/examples/recently_modified.lua | 16 +++ doc/init.lua | 25 +++++ doc/lua.md | 92 ++++++++++++++++ main.go | 4 + pkg/settings/lua/autocmd.go | 165 ++++++++++++++++++++++++++++- pkg/settings/lua/autocmd_test.go | 145 +++++++++++++++++++++++++ pkg/settings/lua/load.go | 9 +- pkg/settings/lua/log.go | 37 +++++++ pkg/settings/lua/log_test.go | 102 ++++++++++++++++++ pkg/settings/lua/lua.go | 13 +++ pkg/upgrade/service.go | 121 +++++++++++++++++++-- pkg/upgrade/service_test.go | 158 +++++++++++++++++++++++++++ sync.go | 1 + 13 files changed, 876 insertions(+), 12 deletions(-) create mode 100644 doc/examples/recently_modified.lua create mode 100644 pkg/settings/lua/log.go create mode 100644 pkg/settings/lua/log_test.go diff --git a/doc/examples/recently_modified.lua b/doc/examples/recently_modified.lua new file mode 100644 index 00000000..1481c5ec --- /dev/null +++ b/doc/examples/recently_modified.lua @@ -0,0 +1,16 @@ +yay.create_autocmd("UpgradeSelect", { + desc = "skip recently modified AUR upgrades", + callback = function(event) + yay.log.info("pre-excluding AUR packages modified in the last 3 days") + local exclude = {} + local recent_cutoff = os.time() - (3 * 24 * 60 * 60) + for _, pkg in ipairs(event.data.upgrades) do + if pkg.repository == "aur" and pkg.last_modified >= recent_cutoff then + yay.log.warn("pre-excluding recently modified AUR package: ", pkg.name) + table.insert(exclude, pkg.name) + end + end + + return { exclude = exclude, skip_menu = false } + end, +}) \ No newline at end of file diff --git a/doc/init.lua b/doc/init.lua index 0e9491b5..f27e4eab 100644 --- a/doc/init.lua +++ b/doc/init.lua @@ -52,7 +52,31 @@ yay.opt.debug = false -- Enable debug logging and local init.lua lookup convenie yay.opt.rpc = true -- Use AUR RPC for dependency/query operations. yay.opt.double_confirm = true -- Ask for confirmation before and after builds during upgrades. +-- Logging +-- yay.log.info("loaded yay init.lua") +-- yay.log.debug("build dir:", yay.opt.build_dir) + -- Hooks +-- Run Lua before yay prints the upgrade exclusion menu. Return package names +-- from event.data.upgrades to pre-exclude them. Set skip_menu = false, or omit +-- it, to show the native menu after these exclusions are applied. +-- +-- yay.create_autocmd("UpgradeSelect", { +-- desc = "skip recently modified AUR upgrades", +-- callback = function(event) +-- local exclude = {} +-- local recent_cutoff = os.time() - (3 * 24 * 60 * 60) +-- for _, pkg in ipairs(event.data.upgrades) do +-- if pkg.repository == "aur" and pkg.last_modified >= recent_cutoff then +-- yay.log.warn("pre-excluding recently modified AUR package:", pkg.name) +-- table.insert(exclude, pkg.name) +-- end +-- end +-- +-- return { exclude = exclude, skip_menu = true } +-- end, +-- }) +-- -- Run Lua after AUR PKGBUILD repos are downloaded/merged and before the -- clean/diff/edit menus or source downloads. -- @@ -60,6 +84,7 @@ yay.opt.double_confirm = true -- Ask for confirmation before and after builds du -- desc = "inspect or modify AUR package files", -- callback = function(event) -- if event.data.pkgbuild:match("forbidden.example") then +-- yay.log.warn(event.match .. ": forbidden source URL") -- yay.abort(event.match .. ": forbidden source URL") -- end -- diff --git a/doc/lua.md b/doc/lua.md index 61a8f31b..fd0b99dd 100644 --- a/doc/lua.md +++ b/doc/lua.md @@ -48,6 +48,97 @@ startup and reports the offending keys/values so misconfigurations fail fast. A ready-to-copy example lives at [`doc/init.lua`](init.lua). +## Logging with `yay.log` + +Lua config and hooks can write through yay's normal logger: + +```lua +yay.log.debug("build dir:", yay.opt.build_dir) +yay.log.info("loaded init.lua") +yay.log.warn("skipping", "pkgname") +yay.log.error("policy check failed") +``` + +`debug` only prints when debug logging is enabled. `error` logs an error-level +message and does not stop execution; use `yay.abort("message")` for controlled +hook stops. + +## Upgrade selection hooks + +`UpgradeSelect` runs during `yay -Syu` after yay has built and sorted the +upgrade graph, and before the native "Packages to exclude" menu is printed. +The hook can return package names to exclude. By default, yay still shows the +native menu after applying hook exclusions. + +```lua +yay.create_autocmd("UpgradeSelect", { + desc = "skip recently modified AUR upgrades", + callback = function(event) + local exclude = {} + local recent_cutoff = os.time() - (3 * 24 * 60 * 60) + for _, pkg in ipairs(event.data.upgrades) do + if pkg.repository == "aur" and pkg.last_modified >= recent_cutoff then + yay.log.debug("pre-excluding recently modified AUR package:", pkg.name) + table.insert(exclude, pkg.name) + end + end + + return { exclude = exclude, skip_menu = true } + end, +}) +``` + +Multiple `UpgradeSelect` hooks run in registration order. Their `exclude` +lists are unioned. If any hook returns `skip_menu = true`, yay applies all hook +exclusions and skips the native menu. With `skip_menu = false` or no return +value, hook exclusions are applied first and then the native menu is shown. + +Returned exclusions must name packages from `event.data.upgrades`. Unknown +names are treated as hook errors so typos do not silently upgrade the wrong +package. Pulled dependencies are visible in `event.data.pulled_dependencies`, +but they are removed only when pruning an excluded upgrade candidate requires +it. + +### UpgradeSelect event + +The callback receives this table: + +```lua +{ + event = "UpgradeSelect", + data = { + upgrades = { + { + id = 3, + name = "pkgname", + base = "pkgbase", + repository = "aur", + local_version = "1.2.3-3", + remote_version = "1.2.3-4", + reason = "explicit", + last_modified = 1700000000, + }, + }, + pulled_dependencies = { + { + id = 0, + name = "depname", + base = "", + repository = "core", + local_version = "", + remote_version = "1.0-1", + reason = "dependency", + last_modified = 0, + }, + }, + }, +} +``` + +For selectable `data.upgrades` entries, `id` matches the number shown in the +native menu. `pulled_dependencies` entries are shown separately by yay and use +`id = 0` because they are not directly selectable. + ## AUR pre-install hooks `init.lua` can register hooks with a small autocmd API: @@ -134,6 +225,7 @@ yay.create_autocmd("AURPreInstall", { desc = "block forbidden sources and patch a PKGBUILD", callback = function(event) if event.data.pkgbuild:match("forbidden.example") then + yay.log.warn(event.match .. ": forbidden source URL") yay.abort(event.match .. ": forbidden source URL") end diff --git a/main.go b/main.go index 82270d98..81c9c9d1 100644 --- a/main.go +++ b/main.go @@ -128,6 +128,10 @@ func main() { return } + + if luaEngine != nil { + luaEngine.SetLogger(run.Logger.Child("lua")) + } run.Lua = luaEngine dbExecutor, err := ialpm.NewExecutor(run.PacmanConf, run.Logger.Child("db")) diff --git a/pkg/settings/lua/autocmd.go b/pkg/settings/lua/autocmd.go index a5505ce2..179544c6 100644 --- a/pkg/settings/lua/autocmd.go +++ b/pkg/settings/lua/autocmd.go @@ -3,10 +3,14 @@ package lua import ( "fmt" + mapset "github.com/deckarep/golang-set/v2" glua "github.com/yuin/gopher-lua" ) -const EventAURPreInstall = "AURPreInstall" +const ( + EventAURPreInstall = "AURPreInstall" + EventUpgradeSelect = "UpgradeSelect" +) type Autocmd struct { Event string @@ -55,9 +59,30 @@ type AURPreInstallSRCINFO struct { Replaces []string } +type UpgradeSelectEvent struct { + Upgrades []UpgradeSelectPackage + PulledDependencies []UpgradeSelectPackage +} + +type UpgradeSelectPackage struct { + ID int + Name string + Base string + Repository string + LocalVersion string + RemoteVersion string + Reason string + LastModified int64 +} + +type UpgradeSelectResult struct { + Exclude []string + SkipMenu bool +} + func (e *Engine) createAutocmd(state *glua.LState) int { event := state.CheckString(1) - if event != EventAURPreInstall { + if event != EventAURPreInstall && event != EventUpgradeSelect { state.ArgError(1, fmt.Sprintf("unsupported event %q", event)) return 0 } @@ -116,6 +141,55 @@ func (e *Engine) RunAURPreInstall(event *AURPreInstallEvent) error { return nil } +func (e *Engine) RunUpgradeSelect(event *UpgradeSelectEvent) (UpgradeSelectResult, error) { + var result UpgradeSelectResult + if !e.HasAutocmd(EventUpgradeSelect) { + return result, nil + } + + validExcludes := mapset.NewThreadUnsafeSetWithSize[string](len(event.Upgrades)) + for _, pkg := range event.Upgrades { + validExcludes.Add(pkg.Name) + } + + seenExcludes := mapset.NewThreadUnsafeSet[string]() + for _, autocmd := range e.autocmds[EventUpgradeSelect] { + if err := e.L.CallByParam(glua.P{ + Fn: autocmd.callback, + NRet: 1, + Protect: true, + }, e.upgradeSelectTable(event)); err != nil { + wrapped := err + if abortErr, ok := luaAbortError(err); ok { + wrapped = abortErr + } + + return result, fmt.Errorf("%s: %w", EventUpgradeSelect, wrapped) + } + + value := e.L.Get(-1) + e.L.Pop(1) + + hookResult, err := e.parseUpgradeSelectResult(value, validExcludes) + if err != nil { + return result, fmt.Errorf("%s: %w", EventUpgradeSelect, err) + } + + for _, name := range hookResult.Exclude { + if !seenExcludes.Add(name) { + continue + } + result.Exclude = append(result.Exclude, name) + } + + if hookResult.SkipMenu { + result.SkipMenu = true + } + } + + return result, nil +} + func (e *Engine) aurPreInstallTable(event *AURPreInstallEvent) *glua.LTable { state := e.L eventTable := state.NewTable() @@ -139,6 +213,20 @@ func (e *Engine) aurPreInstallTable(event *AURPreInstallEvent) *glua.LTable { return eventTable } +func (e *Engine) upgradeSelectTable(event *UpgradeSelectEvent) *glua.LTable { + state := e.L + eventTable := state.NewTable() + data := state.NewTable() + + eventTable.RawSetString("event", glua.LString(EventUpgradeSelect)) + eventTable.RawSetString("data", data) + + data.RawSetString("upgrades", e.upgradeSelectPackagesTable(event.Upgrades)) + data.RawSetString("pulled_dependencies", e.upgradeSelectPackagesTable(event.PulledDependencies)) + + return eventTable +} + func (e *Engine) packagesTable(packages []AURPreInstallPackage) *glua.LTable { state := e.L tbl := state.NewTable() @@ -157,6 +245,26 @@ func (e *Engine) packagesTable(packages []AURPreInstallPackage) *glua.LTable { return tbl } +func (e *Engine) upgradeSelectPackagesTable(packages []UpgradeSelectPackage) *glua.LTable { + state := e.L + tbl := state.NewTable() + + for _, pkg := range packages { + pkgTbl := state.NewTable() + pkgTbl.RawSetString("id", glua.LNumber(pkg.ID)) + pkgTbl.RawSetString("name", glua.LString(pkg.Name)) + pkgTbl.RawSetString("base", glua.LString(pkg.Base)) + pkgTbl.RawSetString("repository", glua.LString(pkg.Repository)) + pkgTbl.RawSetString("local_version", glua.LString(pkg.LocalVersion)) + pkgTbl.RawSetString("remote_version", glua.LString(pkg.RemoteVersion)) + pkgTbl.RawSetString("reason", glua.LString(pkg.Reason)) + pkgTbl.RawSetString("last_modified", glua.LNumber(pkg.LastModified)) + tbl.Append(pkgTbl) + } + + return tbl +} + func (e *Engine) srcinfoTable(srcinfo *AURPreInstallSRCINFO) *glua.LTable { state := e.L tbl := state.NewTable() @@ -181,6 +289,59 @@ func (e *Engine) srcinfoTable(srcinfo *AURPreInstallSRCINFO) *glua.LTable { return tbl } +func (e *Engine) parseUpgradeSelectResult(value glua.LValue, validExcludes mapset.Set[string]) (UpgradeSelectResult, error) { + var result UpgradeSelectResult + if value == glua.LNil { + return result, nil + } + + tbl, ok := value.(*glua.LTable) + if !ok { + return result, fmt.Errorf("callback must return nil or table, got %s", value.Type()) + } + + if excludeValue := tbl.RawGetString("exclude"); excludeValue != glua.LNil { + excludeTbl, ok := excludeValue.(*glua.LTable) + if !ok { + return result, fmt.Errorf("exclude must be a table") + } + + var parseErr error + excludeTbl.ForEach(func(_ glua.LValue, val glua.LValue) { + if parseErr != nil { + return + } + + name, ok := val.(glua.LString) + if !ok { + parseErr = fmt.Errorf("exclude entries must be strings") + return + } + + if !validExcludes.Contains(string(name)) { + parseErr = fmt.Errorf("unknown upgrade exclusion %q", string(name)) + return + } + + result.Exclude = append(result.Exclude, string(name)) + }) + if parseErr != nil { + return result, parseErr + } + } + + if skipMenuValue := tbl.RawGetString("skip_menu"); skipMenuValue != glua.LNil { + skipMenu, ok := skipMenuValue.(glua.LBool) + if !ok { + return result, fmt.Errorf("skip_menu must be a boolean") + } + + result.SkipMenu = bool(skipMenu) + } + + return result, nil +} + func (e *Engine) stringArray(values []string) *glua.LTable { tbl := e.L.NewTable() for _, value := range values { diff --git a/pkg/settings/lua/autocmd_test.go b/pkg/settings/lua/autocmd_test.go index 22d3c477..e5205088 100644 --- a/pkg/settings/lua/autocmd_test.go +++ b/pkg/settings/lua/autocmd_test.go @@ -48,6 +48,23 @@ func TestCreateAutocmdRegistersAndRunsInOrder(t *testing.T) { require.Equal(t, []string{"first:demo-base:demo:demo-base", "second"}, order) } +func TestCreateAutocmdRegistersUpgradeSelect(t *testing.T) { + e := New() + defer e.Close() + + require.NoError(t, e.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + desc = "filter upgrades", + callback = function() end, + }) + `)) + + autocmds := e.autocmds[EventUpgradeSelect] + require.Len(t, autocmds, 1) + require.Equal(t, "filter upgrades", autocmds[0].Desc) + require.True(t, e.HasAutocmd(EventUpgradeSelect)) +} + func TestCreateAutocmdRejectsInvalidEvent(t *testing.T) { e := New() defer e.Close() @@ -107,3 +124,131 @@ func TestRunAURPreInstallReturnsAbortWithoutTraceback(t *testing.T) { err := e.RunAURPreInstall(&AURPreInstallEvent{Base: "demo-base"}) require.EqualError(t, err, "AURPreInstall demo-base: blocked by policy") } + +func TestRunUpgradeSelectEventTableShapeAndReturn(t *testing.T) { + e := New() + defer e.Close() + + require.NoError(t, e.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function(event) + if event.event ~= "UpgradeSelect" then error("bad event") end + if event.data.upgrades[1].id ~= 2 then error("bad upgrade id") end + if event.data.upgrades[1].name ~= "linux" then error("bad upgrade name") end + if event.data.upgrades[1].base ~= "linux" then error("bad upgrade base") end + if event.data.upgrades[1].repository ~= "core" then error("bad upgrade repo") end + if event.data.upgrades[1].local_version ~= "1.0" then error("bad local version") end + if event.data.upgrades[1].remote_version ~= "2.0" then error("bad remote version") end + if event.data.upgrades[1].reason ~= "explicit" then error("bad reason") end + if event.data.upgrades[1].last_modified ~= 123 then error("bad last modified") end + if event.data.upgrades[2].id ~= 1 then error("bad second upgrade id") end + if event.data.pulled_dependencies[1].id ~= 0 then error("bad dependency id") end + if event.data.pulled_dependencies[1].name ~= "new-dep" then error("bad dependency name") end + + return { exclude = { "linux" }, skip_menu = true } + end, + }) + `)) + + result, err := e.RunUpgradeSelect(&UpgradeSelectEvent{ + Upgrades: []UpgradeSelectPackage{ + { + ID: 2, + Name: "linux", + Base: "linux", + Repository: "core", + LocalVersion: "1.0", + RemoteVersion: "2.0", + Reason: "explicit", + LastModified: 123, + }, + {ID: 1, Name: "yay", Base: "yay", Repository: "aur"}, + }, + PulledDependencies: []UpgradeSelectPackage{ + {ID: 0, Name: "new-dep", Repository: "core", Reason: "dependency"}, + }, + }) + require.NoError(t, err) + require.Equal(t, UpgradeSelectResult{Exclude: []string{"linux"}, SkipMenu: true}, result) +} + +func TestRunUpgradeSelectNilReturnMeansNoExclusionsAndNoSkip(t *testing.T) { + e := New() + defer e.Close() + + require.NoError(t, e.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function() end, + }) + `)) + + result, err := e.RunUpgradeSelect(&UpgradeSelectEvent{ + Upgrades: []UpgradeSelectPackage{{Name: "linux"}}, + }) + require.NoError(t, err) + require.Empty(t, result.Exclude) + require.False(t, result.SkipMenu) +} + +func TestRunUpgradeSelectMultipleHooksUnionExclusionsAndSkip(t *testing.T) { + e := New() + defer e.Close() + + require.NoError(t, e.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function() + return { exclude = { "pkg-a" } } + end, + }) + yay.create_autocmd("UpgradeSelect", { + callback = function() + return { exclude = { "pkg-b", "pkg-a" }, skip_menu = true } + end, + }) + `)) + + result, err := e.RunUpgradeSelect(&UpgradeSelectEvent{ + Upgrades: []UpgradeSelectPackage{{Name: "pkg-a"}, {Name: "pkg-b"}}, + }) + require.NoError(t, err) + require.Equal(t, []string{"pkg-a", "pkg-b"}, result.Exclude) + require.True(t, result.SkipMenu) +} + +func TestRunUpgradeSelectRejectsUnknownExcludedPackage(t *testing.T) { + e := New() + defer e.Close() + + require.NoError(t, e.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function() + return { exclude = { "typo" } } + end, + }) + `)) + + _, err := e.RunUpgradeSelect(&UpgradeSelectEvent{ + Upgrades: []UpgradeSelectPackage{{Name: "pkg-a"}}, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "UpgradeSelect") + require.Contains(t, err.Error(), `unknown upgrade exclusion "typo"`) +} + +func TestRunUpgradeSelectReturnsAbortWithoutTraceback(t *testing.T) { + e := New() + defer e.Close() + + require.NoError(t, e.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function() + yay.abort("blocked by policy") + end, + }) + `)) + + _, err := e.RunUpgradeSelect(&UpgradeSelectEvent{ + Upgrades: []UpgradeSelectPackage{{Name: "pkg-a"}}, + }) + require.EqualError(t, err, "UpgradeSelect: blocked by policy") +} diff --git a/pkg/settings/lua/load.go b/pkg/settings/lua/load.go index 46df926f..dbf00cd7 100644 --- a/pkg/settings/lua/load.go +++ b/pkg/settings/lua/load.go @@ -8,8 +8,13 @@ import ( ) // Load loads path, applies its yay.opt values onto cfg, and returns the live engine. -func Load(_ *text.Logger, path string, cfg any) (*Engine, error) { - engine := New() +func Load(logger *text.Logger, path string, cfg any) (*Engine, error) { + var luaLogger *text.Logger + if logger != nil { + luaLogger = logger.Child("lua") + } + + engine := NewWithLogger(luaLogger) if err := engine.L.DoFile(path); err != nil { engine.Close() diff --git a/pkg/settings/lua/log.go b/pkg/settings/lua/log.go new file mode 100644 index 00000000..585a24f7 --- /dev/null +++ b/pkg/settings/lua/log.go @@ -0,0 +1,37 @@ +package lua + +import ( + "github.com/Jguer/yay/v12/pkg/text" + + glua "github.com/yuin/gopher-lua" +) + +const logTableName = "log" + +func (e *Engine) registerLog(yayTbl *glua.LTable) { + logTbl := e.L.NewTable() + e.L.SetField(logTbl, "debug", e.L.NewFunction(e.newLogFn((*text.Logger).Debugln))) + e.L.SetField(logTbl, "info", e.L.NewFunction(e.newLogFn((*text.Logger).Infoln))) + e.L.SetField(logTbl, "warn", e.L.NewFunction(e.newLogFn((*text.Logger).Warnln))) + e.L.SetField(logTbl, "error", e.L.NewFunction(e.newLogFn((*text.Logger).Errorln))) + e.L.SetField(yayTbl, logTableName, logTbl) +} + +func (e *Engine) newLogFn(method func(*text.Logger, ...any)) glua.LGFunction { + return func(state *glua.LState) int { + if e.logger != nil { + method(e.logger, logArgs(state)...) + } + return 0 + } +} + +func logArgs(state *glua.LState) []any { + top := state.GetTop() + args := make([]any, top) + for i := 1; i <= top; i++ { + args[i-1] = state.ToStringMeta(state.Get(i)).String() + } + + return args +} diff --git a/pkg/settings/lua/log_test.go b/pkg/settings/lua/log_test.go new file mode 100644 index 00000000..65898f3c --- /dev/null +++ b/pkg/settings/lua/log_test.go @@ -0,0 +1,102 @@ +package lua + +import ( + "bytes" + "strings" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/Jguer/yay/v12/pkg/text" +) + +func withColorDisabled(t *testing.T) { + t.Helper() + + original := text.UseColor + text.UseColor = false + t.Cleanup(func() { text.UseColor = original }) +} + +func newLogTestEngine(t *testing.T, debug bool) (*Engine, *bytes.Buffer, *bytes.Buffer) { + t.Helper() + withColorDisabled(t) + + var stdout, stderr bytes.Buffer + logger := text.NewLogger(&stdout, &stderr, strings.NewReader(""), debug, "lua") + engine := NewWithLogger(logger) + t.Cleanup(engine.Close) + + return engine, &stdout, &stderr +} + +func TestLogInfoWarnAndErrorUseLoggerStreams(t *testing.T) { + engine, stdout, stderr := newLogTestEngine(t, false) + + require.NoError(t, engine.L.DoString(` + yay.log.info("info message") + yay.log.warn("warn message") + yay.log.error("error message") + `)) + + require.Equal(t, "==> info message\n -> warn message\n", stdout.String()) + require.Equal(t, " -> error message\n", stderr.String()) +} + +func TestLogDebugUsesLoggerDebugGate(t *testing.T) { + engine, stdout, stderr := newLogTestEngine(t, false) + + require.NoError(t, engine.L.DoString(`yay.log.debug("hidden")`)) + require.Empty(t, stdout.String()) + require.Empty(t, stderr.String()) + + debugLogger := text.NewLogger(stdout, stderr, strings.NewReader(""), true, "lua") + engine.SetLogger(debugLogger) + + require.NoError(t, engine.L.DoString(`yay.log.debug("visible")`)) + require.Equal(t, "[DEBUG:lua]visible\n", stdout.String()) + require.Empty(t, stderr.String()) +} + +func TestLogStringifiesLuaArgumentsInOrder(t *testing.T) { + engine, stdout, stderr := newLogTestEngine(t, false) + + require.NoError(t, engine.L.DoString(` + local value = setmetatable({}, { + __tostring = function() + return "custom" + end, + }) + yay.log.info("pkg", 12, true, nil, value) + `)) + + require.Equal(t, "==> pkg 12 true nil custom\n", stdout.String()) + require.Empty(t, stderr.String()) +} + +func TestLogWithoutLoggerDoesNotPanic(t *testing.T) { + engine := New() + t.Cleanup(engine.Close) + + require.NoError(t, engine.L.DoString(` + yay.log.debug("debug") + yay.log.info("info") + yay.log.warn("warn") + yay.log.error("error") + `)) +} + +func TestLoadUsesLuaChildLogger(t *testing.T) { + withColorDisabled(t) + + path := writeLuaFile(t, `yay.log.debug("loaded")`) + var stdout, stderr bytes.Buffer + logger := text.NewLogger(&stdout, &stderr, strings.NewReader(""), true, "fallback") + + engine, err := Load(logger, path, &testConfig{}) + require.NoError(t, err) + t.Cleanup(engine.Close) + + require.Equal(t, "[DEBUG:lua]loaded\n", stdout.String()) + require.Empty(t, stderr.String()) +} diff --git a/pkg/settings/lua/lua.go b/pkg/settings/lua/lua.go index 1039c50e..9f1d5a31 100644 --- a/pkg/settings/lua/lua.go +++ b/pkg/settings/lua/lua.go @@ -5,6 +5,8 @@ import ( "fmt" "reflect" + "github.com/Jguer/yay/v12/pkg/text" + lua "github.com/yuin/gopher-lua" ) @@ -16,13 +18,19 @@ const ( type Engine struct { L *lua.LState autocmds map[string][]Autocmd + logger *text.Logger } func New() *Engine { + return NewWithLogger(nil) +} + +func NewWithLogger(logger *text.Logger) *Engine { state := lua.NewState() engine := &Engine{ L: state, autocmds: make(map[string][]Autocmd), + logger: logger, } yayTbl := state.NewTable() @@ -30,10 +38,15 @@ func New() *Engine { state.SetField(yayTbl, optTableName, state.NewTable()) state.SetField(yayTbl, "abort", state.NewFunction(abort)) state.SetField(yayTbl, "create_autocmd", state.NewFunction(engine.createAutocmd)) + engine.registerLog(yayTbl) return engine } +func (e *Engine) SetLogger(logger *text.Logger) { + e.logger = logger +} + func (e *Engine) Close() { e.L.Close() } diff --git a/pkg/upgrade/service.go b/pkg/upgrade/service.go index f94d8a01..54ce6737 100644 --- a/pkg/upgrade/service.go +++ b/pkg/upgrade/service.go @@ -18,6 +18,7 @@ import ( "github.com/Jguer/yay/v12/pkg/multierror" "github.com/Jguer/yay/v12/pkg/query" "github.com/Jguer/yay/v12/pkg/settings" + settingslua "github.com/Jguer/yay/v12/pkg/settings/lua" "github.com/Jguer/yay/v12/pkg/text" "github.com/Jguer/yay/v12/pkg/vcs" ) @@ -31,6 +32,7 @@ type UpgradeService struct { vcsStore vcs.Store cfg *settings.Configuration log *text.Logger + lua *settingslua.Engine noConfirm bool AURWarnings *query.AURWarnings @@ -52,6 +54,10 @@ func NewUpgradeService(grapher *dep.Grapher, aurCache aur.QueryClient, } } +func (u *UpgradeService) SetLua(engine *settingslua.Engine) { + u.lua = engine +} + // upGraph adds packages to upgrade to the graph. func (u *UpgradeService) upGraph(ctx context.Context, graph *topo.Graph[string, *dep.InstallInfo], enableDowngrade bool, @@ -252,12 +258,7 @@ func (u *UpgradeService) GraphUpgrades(ctx context.Context, return graph, nil } -// userExcludeUpgrades asks the user which packages to exclude from the upgrade and -// removes them from the graph -func (u *UpgradeService) UserExcludeUpgrades(graph *topo.Graph[string, *dep.InstallInfo]) ([]string, error) { - if graph.Len() == 0 { - return []string{}, nil - } +func (u *UpgradeService) upgradeSelection(graph *topo.Graph[string, *dep.InstallInfo]) UpSlice { aurUp, repoUp := u.graphToUpSlice(graph) sort.Sort(repoUp) @@ -280,6 +281,42 @@ func (u *UpgradeService) UserExcludeUpgrades(graph *topo.Graph[string, *dep.Inst } } + return allUp +} + +func (u *UpgradeService) runUpgradeSelectHook(graph *topo.Graph[string, *dep.InstallInfo], allUp UpSlice) ( + excluded []string, skipMenu bool, err error, +) { + if u.lua == nil || !u.lua.HasAutocmd(settingslua.EventUpgradeSelect) { + return []string{}, false, nil + } + + var result settingslua.UpgradeSelectResult + result, err = u.lua.RunUpgradeSelect(upgradeSelectEvent(allUp)) + if err != nil { + return nil, false, err + } + + excluded = u.pruneUpgradeNames(graph, result.Exclude) + + return excluded, result.SkipMenu, nil +} + +func (u *UpgradeService) pruneUpgradeNames(graph *topo.Graph[string, *dep.InstallInfo], names []string) []string { + excluded := make([]string, 0, len(names)) + for _, name := range names { + if !graph.Exists(name) { + continue + } + + u.log.Debugln("pruning", name) + excluded = append(excluded, graph.Prune(name)...) + } + + return excluded +} + +func (u *UpgradeService) printUpgradeSelection(allUp UpSlice) { if len(allUp.PulledDeps) > 0 { u.log.Printf("%s"+text.Bold(" %d ")+"%s\n", text.Bold(text.Cyan("::")), len(allUp.PulledDeps), text.Bold(gotext.Get("%s will also be installed for this operation.", @@ -290,6 +327,33 @@ func (u *UpgradeService) UserExcludeUpgrades(graph *topo.Graph[string, *dep.Inst u.log.Printf("%s"+text.Bold(" %d ")+"%s\n", text.Bold(text.Cyan("::")), len(allUp.Up), text.Bold(gotext.Get("%s to upgrade/install.", gotext.GetN("package", "packages", len(allUp.Up))))) allUp.Print(u.log) +} + +// userExcludeUpgrades asks the user which packages to exclude from the upgrade and +// removes them from the graph +func (u *UpgradeService) UserExcludeUpgrades(graph *topo.Graph[string, *dep.InstallInfo]) ([]string, error) { + if graph.Len() == 0 { + return []string{}, nil + } + allUp := u.upgradeSelection(graph) + + excluded, skipMenu, err := u.runUpgradeSelectHook(graph, allUp) + if err != nil { + return nil, err + } + if skipMenu { + return excluded, nil + } + if len(excluded) > 0 { + // The hook pruned packages; refresh the selection to reflect the + // smaller graph and skip the menu if nothing is left to upgrade. + allUp = u.upgradeSelection(graph) + if len(allUp.Up) == 0 { + return excluded, nil + } + } + + u.printUpgradeSelection(allUp) u.log.Infoln(gotext.Get("Packages to exclude: (eg: \"1 2 3\", \"1-3\", \"^4\" or repo name)")) u.log.Warnln(gotext.Get("Excluding packages may cause partial upgrades and break systems")) @@ -308,10 +372,9 @@ func (u *UpgradeService) UserExcludeUpgrades(graph *topo.Graph[string, *dep.Inst // No exclusions or inclusions specified, return early if noIncludes && len(exclude) == 0 && otherExclude.Cardinality() == 0 { - return []string{}, nil + return excluded, nil } - excluded := make([]string, 0) for i := range allUp.Up { up := &allUp.Up[i] upgradeID := len(allUp.Up) - i @@ -338,3 +401,45 @@ func (u *UpgradeService) UserExcludeUpgrades(graph *topo.Graph[string, *dep.Inst return excluded, nil } + +func upgradeSelectEvent(allUp UpSlice) *settingslua.UpgradeSelectEvent { + return &settingslua.UpgradeSelectEvent{ + Upgrades: upgradeSelectPackages(allUp.Up, true), + PulledDependencies: upgradeSelectPackages(allUp.PulledDeps, false), + } +} + +func upgradeSelectPackages(upgrades []Upgrade, selectable bool) []settingslua.UpgradeSelectPackage { + packages := make([]settingslua.UpgradeSelectPackage, 0, len(upgrades)) + for i := range upgrades { + up := &upgrades[i] + id := 0 + if selectable { + id = len(upgrades) - i + } + + packages = append(packages, settingslua.UpgradeSelectPackage{ + ID: id, + Name: up.Name, + Base: up.Base, + Repository: up.Repository, + LocalVersion: up.LocalVersion, + RemoteVersion: up.RemoteVersion, + Reason: upgradeSelectReason(up.Reason), + LastModified: up.LastModified, + }) + } + + return packages +} + +func upgradeSelectReason(reason alpm.PkgReason) string { + switch reason { + case alpm.PkgReasonExplicit: + return "explicit" + case alpm.PkgReasonDepend: + return "dependency" + default: + return "unknown" + } +} diff --git a/pkg/upgrade/service_test.go b/pkg/upgrade/service_test.go index 67c1be63..c29ff3ab 100644 --- a/pkg/upgrade/service_test.go +++ b/pkg/upgrade/service_test.go @@ -21,6 +21,7 @@ import ( "github.com/Jguer/yay/v12/pkg/dep/topo" "github.com/Jguer/yay/v12/pkg/query" "github.com/Jguer/yay/v12/pkg/settings" + settingslua "github.com/Jguer/yay/v12/pkg/settings/lua" "github.com/Jguer/yay/v12/pkg/settings/parser" "github.com/Jguer/yay/v12/pkg/text" "github.com/Jguer/yay/v12/pkg/vcs" @@ -32,6 +33,77 @@ func ptrString(s string) *string { return &s } +func newUpgradeSelectTestService(input io.Reader, luaEngine *settingslua.Engine) *UpgradeService { + logger := text.NewLogger(io.Discard, io.Discard, input, true, "test") + u := &UpgradeService{ + log: logger, + dbExecutor: &mock.DBExecutor{ + ReposFn: func() []string { return []string{"core"} }, + }, + cfg: &settings.Configuration{Mode: parser.ModeAny}, + AURWarnings: query.NewWarnings(logger), + } + u.SetLua(luaEngine) + + return u +} + +func newUpgradeSelectTestGraph(t *testing.T) *topo.Graph[string, *dep.InstallInfo] { + t.Helper() + + graph := dep.NewGraph() + graph.AddNode("linux") + graph.SetNodeInfo("linux", &topo.NodeInfo[*dep.InstallInfo]{ + Value: &dep.InstallInfo{ + Reason: dep.Explicit, + Source: dep.Sync, + SyncDBName: ptrString("core"), + LocalVersion: "1.0", + Version: "2.0", + Upgrade: true, + }, + }) + + graph.AddNode("yay") + graph.SetNodeInfo("yay", &topo.NodeInfo[*dep.InstallInfo]{ + Value: &dep.InstallInfo{ + Reason: dep.Explicit, + Source: dep.AUR, + AURBase: ptrString("yay"), + LocalVersion: "1.0", + Version: "2.0", + Upgrade: true, + LastModified: 123, + }, + }) + + graph.AddNode("example-git") + graph.AddNode("new-dep") + require.NoError(t, graph.DependOn("example-git", "new-dep")) + graph.SetNodeInfo("example-git", &topo.NodeInfo[*dep.InstallInfo]{ + Value: &dep.InstallInfo{ + Reason: dep.Explicit, + Source: dep.AUR, + AURBase: ptrString("example"), + LocalVersion: "1.0", + Version: "2.0", + Upgrade: true, + LastModified: 456, + }, + }) + graph.SetNodeInfo("new-dep", &topo.NodeInfo[*dep.InstallInfo]{ + Value: &dep.InstallInfo{ + Reason: dep.Dep, + Source: dep.Sync, + SyncDBName: ptrString("core"), + Version: "1.0", + Upgrade: true, + }, + }) + + return graph +} + func TestUpgradeService_GraphUpgrades(t *testing.T) { t.Parallel() linuxDepInfo := &dep.InstallInfo{ @@ -686,6 +758,92 @@ func TestUpgradeService_GraphUpgradesNoUpdates(t *testing.T) { } } +func TestUpgradeService_UserExcludeUpgradesWithoutLuaHookUsesNativeMenu(t *testing.T) { + graph := newUpgradeSelectTestGraph(t) + u := newUpgradeSelectTestService(strings.NewReader("2\n"), nil) + + excluded, err := u.UserExcludeUpgrades(graph) + require.NoError(t, err) + + assert.ElementsMatch(t, []string{"example-git", "new-dep"}, excluded) + assert.True(t, graph.Exists("linux")) + assert.True(t, graph.Exists("yay")) + assert.False(t, graph.Exists("example-git")) + assert.False(t, graph.Exists("new-dep")) +} + +func TestUpgradeService_UserExcludeUpgradesLuaHookPrunesGraph(t *testing.T) { + engine := settingslua.New() + defer engine.Close() + require.NoError(t, engine.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function() + return { exclude = { "example-git" } } + end, + }) + `)) + + graph := newUpgradeSelectTestGraph(t) + u := newUpgradeSelectTestService(strings.NewReader("\n"), engine) + + excluded, err := u.UserExcludeUpgrades(graph) + require.NoError(t, err) + + assert.ElementsMatch(t, []string{"example-git", "new-dep"}, excluded) + assert.True(t, graph.Exists("linux")) + assert.True(t, graph.Exists("yay")) + assert.False(t, graph.Exists("example-git")) + assert.False(t, graph.Exists("new-dep")) +} + +func TestUpgradeService_UserExcludeUpgradesLuaHookSkipMenuAvoidsInput(t *testing.T) { + engine := settingslua.New() + defer engine.Close() + require.NoError(t, engine.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function() + return { exclude = { "yay" }, skip_menu = true } + end, + }) + `)) + + graph := newUpgradeSelectTestGraph(t) + u := newUpgradeSelectTestService(strings.NewReader(""), engine) + + excluded, err := u.UserExcludeUpgrades(graph) + require.NoError(t, err) + + assert.ElementsMatch(t, []string{"yay"}, excluded) + assert.True(t, graph.Exists("linux")) + assert.True(t, graph.Exists("example-git")) + assert.True(t, graph.Exists("new-dep")) + assert.False(t, graph.Exists("yay")) +} + +func TestUpgradeService_UserExcludeUpgradesLuaHookFallsThroughToNativeMenu(t *testing.T) { + engine := settingslua.New() + defer engine.Close() + require.NoError(t, engine.L.DoString(` + yay.create_autocmd("UpgradeSelect", { + callback = function() + return { exclude = { "example-git" }, skip_menu = false } + end, + }) + `)) + + graph := newUpgradeSelectTestGraph(t) + u := newUpgradeSelectTestService(strings.NewReader("1\n"), engine) + + excluded, err := u.UserExcludeUpgrades(graph) + require.NoError(t, err) + + assert.ElementsMatch(t, []string{"example-git", "new-dep", "yay"}, excluded) + assert.True(t, graph.Exists("linux")) + assert.False(t, graph.Exists("yay")) + assert.False(t, graph.Exists("example-git")) + assert.False(t, graph.Exists("new-dep")) +} + func TestUpgradeService_Warnings(t *testing.T) { t.Parallel() dbExe := &mock.DBExecutor{ diff --git a/sync.go b/sync.go index b9d95845..c374b125 100644 --- a/sync.go +++ b/sync.go @@ -58,6 +58,7 @@ func syncInstall(ctx context.Context, upService := upgrade.NewUpgradeService( grapher, aurCache, dbExecutor, run.VCSStore, run.Cfg, settings.NoConfirm, run.Logger.Child("upgrade")) + upService.SetLua(run.Lua) graph, errSysUp = upService.GraphUpgrades(ctx, graph, cmdArgs.ExistsDouble("u", "sysupgrade"),