diff --git a/internal/extension/manager.go b/internal/extension/manager.go index 9d9c8542..0748d76a 100644 --- a/internal/extension/manager.go +++ b/internal/extension/manager.go @@ -1,13 +1,8 @@ package extension import ( - "context" - "fmt" "log/slog" - "os" - "path/filepath" "sort" - "strings" "sync" "sync/atomic" "time" @@ -104,155 +99,6 @@ func NewManager(logger *slog.Logger) *Manager { } } -// LoadPaths discovers and loads Lua extension sources from files or directories. -func (manager *Manager) LoadPaths(ctx context.Context, paths []string) error { - for _, extensionPath := range paths { - if err := ctx.Err(); err != nil { - return err - } - if err := manager.loadPath(ctx, extensionPath); err != nil { - return err - } - } - - return nil -} - -func (manager *Manager) loadPath(ctx context.Context, extensionPath string) error { - sources, err := discoverLuaSources(extensionPath) - if err != nil { - return err - } - - for _, source := range sources { - if err := manager.loadSource(ctx, source); err != nil { - return err - } - } - - return nil -} - -func (manager *Manager) loadSource(ctx context.Context, source luaSource) error { - if source.Manifest { - return manager.LoadManifest(ctx, source.Path) - } - - return manager.LoadFile(ctx, source.Path) -} - -// LoadFile loads one Lua extension source file. -func (manager *Manager) LoadFile(ctx context.Context, extensionPath string) error { - return manager.loadLuaFile(ctx, extensionPath, extensionName(extensionPath), extensionPath) -} - -// LoadManifest loads one directory-based Lua extension manifest. -func (manager *Manager) LoadManifest(ctx context.Context, manifestPath string) error { - manifest, err := manager.ReadManifest(manifestPath) - if err != nil { - return err - } - - entry := strings.TrimSpace(manifest.Entry) - if entry == "" { - entry = "main.lua" - } - if filepath.IsAbs(entry) || strings.Contains(entry, "..") { - return fmt.Errorf("extension: invalid entry %q", manifest.Entry) - } - - entryPath := filepath.Join(filepath.Dir(manifestPath), entry) - name := strings.TrimSpace(manifest.Name) - if name == "" { - name = extensionName(filepath.Dir(manifestPath)) - } - - return manager.loadLuaFile(ctx, entryPath, name, filepath.Dir(manifestPath)) -} - -// ReadManifest reads a directory-based Lua manifest without executing extension entry code. -func (manager *Manager) ReadManifest(manifestPath string) (Manifest, error) { - absolutePath, err := filepath.Abs(manifestPath) - if err != nil { - return Manifest{}, fmt.Errorf("extension: resolve manifest: %w", err) - } - - state := lua.NewState(lua.Options{SkipOpenLibs: true}) - defer state.Close() - openExtensionLibs(state) - - if err := state.DoFile(absolutePath); err != nil { - return Manifest{}, fmt.Errorf("extension: load manifest %s: %w", absolutePath, err) - } - table, ok := state.Get(-1).(*lua.LTable) - if !ok { - return Manifest{}, fmt.Errorf("extension: manifest %s must return a table", absolutePath) - } - - manifest := Manifest{ - Name: luaTableString(table, "name", ""), - Version: luaTableString(table, "version", ""), - APIVersion: luaTableString(table, "api_version", ""), - Description: luaTableString(table, "description", ""), - Entry: luaTableString(table, "entry", ""), - } - if strings.TrimSpace(manifest.Name) == "" { - manifest.Name = extensionName(filepath.Dir(absolutePath)) - } - - return manifest, nil -} - -func (manager *Manager) loadLuaFile(ctx context.Context, extensionPath, name, displayPath string) error { - if err := ctx.Err(); err != nil { - return err - } - - absolutePath, err := filepath.Abs(extensionPath) - if err != nil { - return fmt.Errorf("extension: resolve path: %w", err) - } - - manager.addModuleRootsForPath(absolutePath) - extensionRuntime := &luaExtension{ - activeEvent: nil, - state: lua.NewState(lua.Options{SkipOpenLibs: true}), - name: name, - path: displayPath, - commands: []string{}, - tools: []string{}, - keymaps: []string{}, - handlers: []string{}, - lock: sync.Mutex{}, - totalDuration: atomic.Int64{}, - } - openExtensionLibs(extensionRuntime.state) - manager.configurePackagePath(extensionRuntime.state) - manager.installAPI(extensionRuntime) - - startedAt := time.Now() - if err := extensionRuntime.state.DoFile(absolutePath); err != nil { - extensionRuntime.state.Close() - return fmt.Errorf("extension: load %s: %w", absolutePath, err) - } - if setupFn, ok := extensionRuntime.state.Get(-1).(*lua.LFunction); ok { - extensionRuntime.state.Push(setupFn) - extensionRuntime.state.Push(extensionRuntime.state.GetGlobal("librecode")) - if err := extensionRuntime.state.PCall(1, 0, nil); err != nil { - extensionRuntime.state.Close() - return fmt.Errorf("extension: setup %s: %w", absolutePath, err) - } - } - recordLuaCallDuration(extensionRuntime, startedAt) - - manager.lock.Lock() - manager.extensions = append(manager.extensions, extensionRuntime) - manager.lock.Unlock() - manager.logger.Debug("loaded lua extension", slog.String("path", absolutePath)) - - return nil -} - // Extensions returns loaded extension metadata. func (manager *Manager) Extensions() []LoadedExtension { manager.lock.RLock() @@ -307,133 +153,6 @@ func (manager *Manager) Tools() []Tool { return tools } -// ExecuteCommand runs a registered extension slash command. -func (manager *Manager) ExecuteCommand(ctx context.Context, name, args string) (string, error) { - manager.lock.RLock() - command, ok := manager.commands[name] - manager.lock.RUnlock() - if !ok { - return "", fmt.Errorf("extension: command %q not found", name) - } - - result, err := callLua(command.extension, command.function, lua.LString(args)) - if err != nil { - return "", fmt.Errorf("extension: command %q failed: %w", name, err) - } - - if err := ctx.Err(); err != nil { - return "", err - } - - return result.String(), nil -} - -// ExecuteTool runs a registered extension tool. -func (manager *Manager) ExecuteTool(ctx context.Context, name string, args map[string]any) (ToolResult, error) { - manager.lock.RLock() - tool, ok := manager.tools[name] - manager.lock.RUnlock() - if !ok { - return ToolResult{Details: map[string]any{}, Content: ""}, fmt.Errorf("extension: tool %q not found", name) - } - - result, err := callLuaPrepared(tool.extension, nil, tool.function, func(state *lua.LState) []lua.LValue { - return []lua.LValue{mapToLuaTable(state, args)} - }) - if err != nil { - return ToolResult{Details: map[string]any{}, Content: ""}, - fmt.Errorf("extension: tool %q failed: %w", name, err) - } - if err := ctx.Err(); err != nil { - return ToolResult{Details: map[string]any{}, Content: ""}, err - } - - return luaToolResult(result), nil -} - -// HandleTerminalEvent runs registered low-level terminal runtime handlers. -func (manager *Manager) HandleTerminalEvent(ctx context.Context, event *TerminalEvent) (TerminalEventResult, error) { - hostEvent := newLuaHostEvent(event) - if err := manager.runDueTimers(ctx, hostEvent, time.Now()); err != nil { - return hostEvent.result(), err - } - if event.Name == luaFieldKey { - if err := manager.runKeymaps(ctx, hostEvent); err != nil { - return hostEvent.result(), err - } - if hostEvent.stopped { - return hostEvent.result(), nil - } - } - - for _, handler := range manager.handlersFor(event.Name) { - if err := ctx.Err(); err != nil { - return hostEvent.result(), err - } - - result, err := callLuaPrepared( - handler.extension, - hostEvent, - handler.function, - func(state *lua.LState) []lua.LValue { - return []lua.LValue{terminalEventTable(state, hostEvent.eventSnapshot())} - }, - ) - if err != nil { - return hostEvent.result(), fmt.Errorf("extension: terminal event %q failed: %w", event.Name, err) - } - hostEvent.applyLuaResult(result) - if hostEvent.stopped { - break - } - } - - return hostEvent.result(), nil -} - -// Emit sends an event to registered extension handlers. -func (manager *Manager) Emit(ctx context.Context, eventName string, payload map[string]any) error { - for _, handler := range manager.handlersFor(eventName) { - if err := ctx.Err(); err != nil { - return err - } - - _, err := callLuaPrepared(handler.extension, nil, handler.function, func(state *lua.LState) []lua.LValue { - return []lua.LValue{lua.LString(eventName), mapToLuaTable(state, payload)} - }) - if err != nil { - return fmt.Errorf("extension: event %q failed: %w", eventName, err) - } - } - - return nil -} - -// HasTerminalEventHandlers reports whether any extension handler is registered for eventName. -func (manager *Manager) HasTerminalEventHandlers(eventName string) bool { - manager.lock.RLock() - defer manager.lock.RUnlock() - - return len(manager.handlers[eventName]) > 0 || eventName == luaFieldKey && len(manager.keymaps) > 0 -} - -func (manager *Manager) handlersFor(eventName string) []luaHookHandler { - manager.lock.RLock() - handlers := append([]luaHookHandler{}, manager.handlers[eventName]...) - manager.lock.RUnlock() - sort.SliceStable(handlers, func(leftIndex, rightIndex int) bool { - left := handlers[leftIndex] - right := handlers[rightIndex] - if left.priority == right.priority { - return left.order < right.order - } - - return left.priority > right.priority - }) - - return handlers -} - // Shutdown closes all loaded Lua states and clears registrations. func (manager *Manager) Shutdown() { manager.lock.Lock() @@ -456,324 +175,88 @@ func (manager *Manager) Shutdown() { manager.nextNamespaceID = 1 } -func (manager *Manager) installAPI(extensionRuntime *luaExtension) { - apiTable := extensionRuntime.state.NewTable() - extensionRuntime.state.SetFuncs(apiTable, map[string]lua.LGFunction{ - "register_command": manager.luaRegisterCommand(extensionRuntime), - "register_tool": manager.luaRegisterTool(extensionRuntime), - "on": manager.luaOn(extensionRuntime), - "log": manager.luaLog(extensionRuntime), - }) - extensionRuntime.state.SetField(apiTable, "api", manager.luaCoreAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "autocmd", manager.luaAutocmdAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "buf", manager.luaBufferAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "command", manager.luaCommandAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "event", manager.luaEventAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "action", manager.luaActionAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "timer", manager.luaTimerAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "keymap", manager.luaKeymapAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "layout", manager.luaLayoutAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "ui", manager.luaUIAPI(extensionRuntime)) - extensionRuntime.state.SetField(apiTable, "win", manager.luaWindowAPI(extensionRuntime)) - extensionRuntime.state.SetGlobal("librecode", apiTable) - extensionRuntime.state.PreloadModule("librecode", func(state *lua.LState) int { - state.Push(apiTable) - - return 1 - }) -} - -func (manager *Manager) luaRegisterCommand(extensionRuntime *luaExtension) lua.LGFunction { - return func(state *lua.LState) int { - name, description, function := luaRegistrationArgs(state) - definition := Command{Name: name, Description: description, Extension: extensionRuntime.name} - - manager.lock.Lock() - manager.commands[name] = luaCommand{extension: extensionRuntime, function: function, definition: definition} - extensionRuntime.commands = append(extensionRuntime.commands, name) - manager.lock.Unlock() - - return 0 - } -} - -func (manager *Manager) luaRegisterTool(extensionRuntime *luaExtension) lua.LGFunction { - return func(state *lua.LState) int { - name, description, function := luaRegistrationArgs(state) - definition := Tool{ - InputSchema: luaOptionalSchema(state, 4), - Name: name, - Description: description, - Extension: extensionRuntime.name, - } - - manager.lock.Lock() - manager.tools[name] = luaTool{extension: extensionRuntime, function: function, definition: definition} - extensionRuntime.tools = append(extensionRuntime.tools, name) - manager.lock.Unlock() - - return 0 - } -} - -func (manager *Manager) luaOn(extensionRuntime *luaExtension) lua.LGFunction { - return func(state *lua.LState) int { - eventName := state.CheckString(1) - priority, function := luaEventHandlerArgs(state) - manager.registerHandler(extensionRuntime, eventName, priority, function) - - return 0 - } -} - -func (manager *Manager) luaLog(extensionRuntime *luaExtension) lua.LGFunction { - return func(state *lua.LState) int { - message := state.CheckString(1) - manager.logger.Info( - "lua extension", - slog.String("extension", extensionRuntime.name), - slog.String("message", message), - ) - - return 0 - } -} - -func luaRegistrationArgs(state *lua.LState) (name, description string, function *lua.LFunction) { - return state.CheckString(1), state.OptString(2, ""), state.CheckFunction(3) -} - -func luaOptionalSchema(state *lua.LState, index int) map[string]any { - if state.GetTop() < index { - return map[string]any{} - } - if table, ok := state.Get(index).(*lua.LTable); ok { - return luaTableToMap(table) - } - - return map[string]any{} -} - -func luaEventHandlerArgs(state *lua.LState) (priority int, function *lua.LFunction) { - if handler, ok := state.Get(2).(*lua.LFunction); ok { - return 0, handler - } - - options := state.CheckTable(2) - - return int(lua.LVAsNumber(options.RawGetString("priority"))), state.CheckFunction(3) -} - -func callLua(extensionRuntime *luaExtension, function *lua.LFunction, args ...lua.LValue) (lua.LValue, error) { - return callLuaPrepared(extensionRuntime, nil, function, func(*lua.LState) []lua.LValue { - return args - }) -} - -func callLuaPrepared( - extensionRuntime *luaExtension, - hostEvent *luaHostEvent, - function *lua.LFunction, - prepareArgs func(*lua.LState) []lua.LValue, -) (lua.LValue, error) { - extensionRuntime.lock.Lock() - defer extensionRuntime.lock.Unlock() - - previousEvent := extensionRuntime.activeEvent - extensionRuntime.activeEvent = hostEvent - defer func() { - extensionRuntime.activeEvent = previousEvent - }() - - top := extensionRuntime.state.GetTop() - startedAt := time.Now() - defer recordLuaCallDuration(extensionRuntime, startedAt) - - args := prepareArgs(extensionRuntime.state) - if err := extensionRuntime.state.CallByParam(lua.P{Fn: function, NRet: 1, Protect: true}, args...); err != nil { - extensionRuntime.state.SetTop(top) - return lua.LNil, err - } - - result := extensionRuntime.state.Get(-1) - extensionRuntime.state.Pop(1) - extensionRuntime.state.SetTop(top) - - return result, nil -} +func (manager *Manager) unregisterRuntime(extensionRuntime *luaExtension) { + manager.lock.Lock() + defer manager.lock.Unlock() -func recordLuaCallDuration(extensionRuntime *luaExtension, startedAt time.Time) { - extensionRuntime.totalDuration.Add(int64(time.Since(startedAt))) + manager.unregisterRuntimeLocked(extensionRuntime) } -type luaSource struct { - Path string - Manifest bool +func (manager *Manager) unregisterRuntimeLocked(extensionRuntime *luaExtension) { + manager.unregisterCommandsLocked(extensionRuntime) + manager.unregisterToolsLocked(extensionRuntime) + manager.unregisterHandlersLocked(extensionRuntime) + manager.unregisterKeymapsLocked(extensionRuntime) + manager.unregisterTimersLocked(extensionRuntime) + manager.unregisterExtensionLocked(extensionRuntime) } -func discoverLuaSources(extensionPath string) ([]luaSource, error) { - if extensionPath == "" { - return []luaSource{}, nil - } - - info, err := os.Stat(extensionPath) - if err != nil { - if os.IsNotExist(err) { - return []luaSource{}, nil - } - return nil, fmt.Errorf("extension: stat %s: %w", extensionPath, err) - } - - if !info.IsDir() { - if strings.HasSuffix(extensionPath, ".lua") { - return []luaSource{{Path: extensionPath, Manifest: false}}, nil +func (manager *Manager) unregisterCommandsLocked(extensionRuntime *luaExtension) { + for _, name := range extensionRuntime.commands { + if command, ok := manager.commands[name]; ok && command.extension == extensionRuntime { + delete(manager.commands, name) } - return []luaSource{}, nil } - - return discoverLuaDir(extensionPath) } -func discoverLuaDir(root string) ([]luaSource, error) { - manifestPath := filepath.Join(root, "init.lua") - if info, err := os.Stat(manifestPath); err == nil && !info.IsDir() { - return []luaSource{{Path: manifestPath, Manifest: true}}, nil - } - - sources := []luaSource{} - walkErr := filepath.WalkDir(root, func(currentPath string, dirEntry os.DirEntry, walkErr error) error { - return collectLuaSource(root, currentPath, dirEntry, walkErr, &sources) - }) - if walkErr != nil { - return nil, fmt.Errorf("extension: walk %s: %w", root, walkErr) +func (manager *Manager) unregisterToolsLocked(extensionRuntime *luaExtension) { + for _, name := range extensionRuntime.tools { + if tool, ok := manager.tools[name]; ok && tool.extension == extensionRuntime { + delete(manager.tools, name) + } } - sort.Slice(sources, func(i, j int) bool { return sources[i].Path < sources[j].Path }) - - return sources, nil } -func collectLuaSource(root, currentPath string, dirEntry os.DirEntry, walkErr error, sources *[]luaSource) error { - if walkErr != nil { - return walkErr - } - if isExtensionDir(root, currentPath, dirEntry) { - return collectLuaSourceDir(root, currentPath, sources) - } - if strings.HasSuffix(currentPath, ".lua") { - *sources = append(*sources, luaSource{Path: currentPath, Manifest: false}) +func (manager *Manager) unregisterHandlersLocked(extensionRuntime *luaExtension) { + for eventName, handlers := range manager.handlers { + filtered := keepHandlersFromOtherRuntimes(handlers, extensionRuntime) + if len(filtered) == 0 { + delete(manager.handlers, eventName) + continue + } + manager.handlers[eventName] = filtered } - - return nil } -func isExtensionDir(root, currentPath string, dirEntry os.DirEntry) bool { - if dirEntry.IsDir() { - return true - } - if currentPath == root || dirEntry.Type()&os.ModeSymlink == 0 { - return false +func keepHandlersFromOtherRuntimes(handlers []luaHookHandler, extensionRuntime *luaExtension) []luaHookHandler { + filtered := handlers[:0] + for _, handler := range handlers { + if handler.extension != extensionRuntime { + filtered = append(filtered, handler) + } } - info, err := os.Stat(currentPath) - return err == nil && info.IsDir() + return filtered } -func collectLuaSourceDir(root, currentPath string, sources *[]luaSource) error { - if currentPath == root { - return nil - } - if filepath.Base(currentPath) == "lua" { - return filepath.SkipDir - } - manifestPath := filepath.Join(currentPath, "init.lua") - if info, err := os.Stat(manifestPath); err == nil && !info.IsDir() { - *sources = append(*sources, luaSource{Path: manifestPath, Manifest: true}) - return filepath.SkipDir +func (manager *Manager) unregisterKeymapsLocked(extensionRuntime *luaExtension) { + filtered := manager.keymaps[:0] + for _, keymap := range manager.keymaps { + if keymap.extension != extensionRuntime { + filtered = append(filtered, keymap) + } } - - return nil + manager.keymaps = filtered } -func (manager *Manager) addModuleRootsForPath(extensionPath string) { - if strings.TrimSpace(extensionPath) == "" { - return - } - absolutePath, err := filepath.Abs(extensionPath) - if err != nil { - absolutePath = extensionPath - } - roots := moduleRootsForPath(absolutePath) - - manager.lock.Lock() - defer manager.lock.Unlock() - - seen := make(map[string]struct{}, len(manager.moduleRoots)+len(roots)) - for _, root := range manager.moduleRoots { - seen[root] = struct{}{} - } - for _, root := range roots { - if _, ok := seen[root]; ok { +func (manager *Manager) unregisterTimersLocked(extensionRuntime *luaExtension) { + filtered := manager.timers[:0] + for _, timer := range manager.timers { + if timer.extension != extensionRuntime { + filtered = append(filtered, timer) continue } - manager.moduleRoots = append(manager.moduleRoots, root) - seen[root] = struct{}{} - } -} - -func moduleRootsForPath(extensionPath string) []string { - root := extensionPath - if info, err := os.Stat(extensionPath); err == nil && !info.IsDir() { - root = filepath.Dir(extensionPath) - } - - return []string{root} -} - -func (manager *Manager) configurePackagePath(state *lua.LState) { - packageTable, ok := state.GetGlobal("package").(*lua.LTable) - if !ok { - return - } - patterns := []string{packageTable.RawGetString("path").String()} - for _, root := range manager.moduleRootsSnapshot() { - patterns = append(patterns, - filepath.ToSlash(filepath.Join(root, "?.lua")), - filepath.ToSlash(filepath.Join(root, "?", "init.lua")), - ) + manager.canceledTimers[timer.id] = struct{}{} } - packageTable.RawSetString("path", lua.LString(strings.Join(patterns, ";"))) -} - -func (manager *Manager) moduleRootsSnapshot() []string { - manager.lock.RLock() - defer manager.lock.RUnlock() - - return append([]string{}, manager.moduleRoots...) + manager.timers = filtered } -func openExtensionLibs(state *lua.LState) { - libraries := []struct { - open lua.LGFunction - name string - }{ - {name: lua.BaseLibName, open: lua.OpenBase}, - {name: lua.LoadLibName, open: lua.OpenPackage}, - {name: lua.TabLibName, open: lua.OpenTable}, - {name: lua.StringLibName, open: lua.OpenString}, - {name: lua.MathLibName, open: lua.OpenMath}, - {name: lua.IoLibName, open: lua.OpenIo}, - {name: lua.OsLibName, open: lua.OpenOs}, - {name: lua.DebugLibName, open: lua.OpenDebug}, - } - - for _, library := range libraries { - state.Push(state.NewFunction(library.open)) - state.Push(lua.LString(library.name)) - state.Call(1, 0) +func (manager *Manager) unregisterExtensionLocked(extensionRuntime *luaExtension) { + filtered := manager.extensions[:0] + for _, loadedRuntime := range manager.extensions { + if loadedRuntime != extensionRuntime { + filtered = append(filtered, loadedRuntime) + } } -} - -func extensionName(extensionPath string) string { - baseName := filepath.Base(extensionPath) - return strings.TrimSuffix(baseName, filepath.Ext(baseName)) + manager.extensions = filtered } diff --git a/internal/extension/manager_dispatch.go b/internal/extension/manager_dispatch.go new file mode 100644 index 00000000..697d10ca --- /dev/null +++ b/internal/extension/manager_dispatch.go @@ -0,0 +1,187 @@ +package extension + +import ( + "context" + "fmt" + "sort" + "time" + + lua "github.com/yuin/gopher-lua" +) + +// ExecuteCommand runs a registered extension slash command. +func (manager *Manager) ExecuteCommand(ctx context.Context, name, args string) (string, error) { + manager.lock.RLock() + command, ok := manager.commands[name] + manager.lock.RUnlock() + if !ok { + return "", fmt.Errorf("extension: command %q not found", name) + } + + if err := ctx.Err(); err != nil { + return "", err + } + + result, err := callLua(command.extension, command.function, lua.LString(args)) + if err != nil { + return "", fmt.Errorf("extension: command %q failed: %w", name, err) + } + + if err := ctx.Err(); err != nil { + return "", err + } + + return result.String(), nil +} + +// ExecuteTool runs a registered extension tool. +func (manager *Manager) ExecuteTool(ctx context.Context, name string, args map[string]any) (ToolResult, error) { + manager.lock.RLock() + tool, ok := manager.tools[name] + manager.lock.RUnlock() + if !ok { + return ToolResult{Details: map[string]any{}, Content: ""}, fmt.Errorf("extension: tool %q not found", name) + } + + if err := ctx.Err(); err != nil { + return ToolResult{Details: map[string]any{}, Content: ""}, err + } + + result, err := callLuaPrepared(tool.extension, nil, tool.function, func(state *lua.LState) []lua.LValue { + return []lua.LValue{mapToLuaTable(state, args)} + }) + if err != nil { + return ToolResult{Details: map[string]any{}, Content: ""}, + fmt.Errorf("extension: tool %q failed: %w", name, err) + } + if err := ctx.Err(); err != nil { + return ToolResult{Details: map[string]any{}, Content: ""}, err + } + + return luaToolResult(result), nil +} + +// HandleTerminalEvent runs registered low-level terminal runtime handlers. +func (manager *Manager) HandleTerminalEvent(ctx context.Context, event *TerminalEvent) (TerminalEventResult, error) { + hostEvent := newLuaHostEvent(event) + if err := manager.runDueTimers(ctx, hostEvent, time.Now()); err != nil { + return hostEvent.result(), err + } + if event.Name == luaFieldKey { + if err := manager.runKeymaps(ctx, hostEvent); err != nil { + return hostEvent.result(), err + } + if hostEvent.stopped { + return hostEvent.result(), nil + } + } + + for _, handler := range manager.handlersFor(event.Name) { + if err := ctx.Err(); err != nil { + return hostEvent.result(), err + } + + result, err := callLuaPrepared( + handler.extension, + hostEvent, + handler.function, + func(state *lua.LState) []lua.LValue { + return []lua.LValue{terminalEventTable(state, hostEvent.eventSnapshot())} + }, + ) + if err != nil { + return hostEvent.result(), fmt.Errorf("extension: terminal event %q failed: %w", event.Name, err) + } + hostEvent.applyLuaResult(result) + if hostEvent.stopped { + break + } + } + + return hostEvent.result(), nil +} + +// Emit sends an event to registered extension handlers. +func (manager *Manager) Emit(ctx context.Context, eventName string, payload map[string]any) error { + for _, handler := range manager.handlersFor(eventName) { + if err := ctx.Err(); err != nil { + return err + } + + _, err := callLuaPrepared(handler.extension, nil, handler.function, func(state *lua.LState) []lua.LValue { + return []lua.LValue{lua.LString(eventName), mapToLuaTable(state, payload)} + }) + if err != nil { + return fmt.Errorf("extension: event %q failed: %w", eventName, err) + } + } + + return nil +} + +// HasTerminalEventHandlers reports whether any extension handler is registered for eventName. +func (manager *Manager) HasTerminalEventHandlers(eventName string) bool { + manager.lock.RLock() + defer manager.lock.RUnlock() + + return len(manager.handlers[eventName]) > 0 || eventName == luaFieldKey && len(manager.keymaps) > 0 +} + +func (manager *Manager) handlersFor(eventName string) []luaHookHandler { + manager.lock.RLock() + handlers := append([]luaHookHandler{}, manager.handlers[eventName]...) + manager.lock.RUnlock() + sort.SliceStable(handlers, func(leftIndex, rightIndex int) bool { + left := handlers[leftIndex] + right := handlers[rightIndex] + if left.priority == right.priority { + return left.order < right.order + } + + return left.priority > right.priority + }) + + return handlers +} + +func callLua(extensionRuntime *luaExtension, function *lua.LFunction, args ...lua.LValue) (lua.LValue, error) { + return callLuaPrepared(extensionRuntime, nil, function, func(*lua.LState) []lua.LValue { + return args + }) +} + +func callLuaPrepared( + extensionRuntime *luaExtension, + hostEvent *luaHostEvent, + function *lua.LFunction, + prepareArgs func(*lua.LState) []lua.LValue, +) (lua.LValue, error) { + extensionRuntime.lock.Lock() + defer extensionRuntime.lock.Unlock() + + previousEvent := extensionRuntime.activeEvent + extensionRuntime.activeEvent = hostEvent + defer func() { + extensionRuntime.activeEvent = previousEvent + }() + + top := extensionRuntime.state.GetTop() + startedAt := time.Now() + defer recordLuaCallDuration(extensionRuntime, startedAt) + + args := prepareArgs(extensionRuntime.state) + if err := extensionRuntime.state.CallByParam(lua.P{Fn: function, NRet: 1, Protect: true}, args...); err != nil { + extensionRuntime.state.SetTop(top) + return lua.LNil, err + } + + result := extensionRuntime.state.Get(-1) + extensionRuntime.state.Pop(1) + extensionRuntime.state.SetTop(top) + + return result, nil +} + +func recordLuaCallDuration(extensionRuntime *luaExtension, startedAt time.Time) { + extensionRuntime.totalDuration.Add(int64(time.Since(startedAt))) +} diff --git a/internal/extension/manager_loader.go b/internal/extension/manager_loader.go new file mode 100644 index 00000000..cf537475 --- /dev/null +++ b/internal/extension/manager_loader.go @@ -0,0 +1,300 @@ +package extension + +import ( + "context" + "fmt" + "log/slog" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "time" + + lua "github.com/yuin/gopher-lua" +) + +// LoadPaths discovers and loads Lua extension sources from files or directories. +func (manager *Manager) LoadPaths(ctx context.Context, paths []string) error { + for _, extensionPath := range paths { + if err := ctx.Err(); err != nil { + return err + } + if err := manager.loadPath(ctx, extensionPath); err != nil { + return err + } + } + + return nil +} + +func (manager *Manager) loadPath(ctx context.Context, extensionPath string) error { + sources, err := discoverLuaSources(extensionPath) + if err != nil { + return err + } + + for _, source := range sources { + if err := manager.loadSource(ctx, source); err != nil { + return err + } + } + + return nil +} + +func (manager *Manager) loadSource(ctx context.Context, source luaSource) error { + if source.Manifest { + return manager.LoadManifest(ctx, source.Path) + } + + return manager.LoadFile(ctx, source.Path) +} + +// LoadFile loads one Lua extension source file. +func (manager *Manager) LoadFile(ctx context.Context, extensionPath string) error { + return manager.loadLuaFile(ctx, extensionPath, extensionName(extensionPath), extensionPath) +} + +// LoadManifest loads one directory-based Lua extension manifest. +func (manager *Manager) LoadManifest(ctx context.Context, manifestPath string) error { + manifest, err := manager.ReadManifest(manifestPath) + if err != nil { + return err + } + + entry := strings.TrimSpace(manifest.Entry) + if entry == "" { + entry = "main.lua" + } + if filepath.IsAbs(entry) || strings.Contains(entry, "..") { + return fmt.Errorf("extension: invalid entry %q", manifest.Entry) + } + + entryPath := filepath.Join(filepath.Dir(manifestPath), entry) + name := strings.TrimSpace(manifest.Name) + if name == "" { + name = extensionName(filepath.Dir(manifestPath)) + } + + return manager.loadLuaFile(ctx, entryPath, name, filepath.Dir(manifestPath)) +} + +// ReadManifest reads a directory-based Lua manifest without executing extension entry code. +func (manager *Manager) ReadManifest(manifestPath string) (Manifest, error) { + absolutePath, err := filepath.Abs(manifestPath) + if err != nil { + return Manifest{}, fmt.Errorf("extension: resolve manifest: %w", err) + } + + state := lua.NewState(lua.Options{SkipOpenLibs: true}) + defer state.Close() + openExtensionLibs(state) + + if err := state.DoFile(absolutePath); err != nil { + return Manifest{}, fmt.Errorf("extension: load manifest %s: %w", absolutePath, err) + } + table, ok := state.Get(-1).(*lua.LTable) + if !ok { + return Manifest{}, fmt.Errorf("extension: manifest %s must return a table", absolutePath) + } + + manifest := Manifest{ + Name: luaTableString(table, "name", ""), + Version: luaTableString(table, "version", ""), + APIVersion: luaTableString(table, "api_version", ""), + Description: luaTableString(table, "description", ""), + Entry: luaTableString(table, "entry", ""), + } + if strings.TrimSpace(manifest.Name) == "" { + manifest.Name = extensionName(filepath.Dir(absolutePath)) + } + + return manifest, nil +} + +func (manager *Manager) loadLuaFile(ctx context.Context, extensionPath, name, displayPath string) error { + if err := ctx.Err(); err != nil { + return err + } + + absolutePath, err := filepath.Abs(extensionPath) + if err != nil { + return fmt.Errorf("extension: resolve path: %w", err) + } + + manager.addModuleRootsForPath(absolutePath) + extensionRuntime := &luaExtension{ + activeEvent: nil, + state: lua.NewState(lua.Options{SkipOpenLibs: true}), + name: name, + path: displayPath, + commands: []string{}, + tools: []string{}, + keymaps: []string{}, + handlers: []string{}, + lock: sync.Mutex{}, + totalDuration: atomic.Int64{}, + } + openExtensionLibs(extensionRuntime.state) + manager.configurePackagePath(extensionRuntime.state) + manager.installAPI(extensionRuntime) + + startedAt := time.Now() + if err := extensionRuntime.state.DoFile(absolutePath); err != nil { + manager.unregisterRuntime(extensionRuntime) + extensionRuntime.state.Close() + return fmt.Errorf("extension: load %s: %w", absolutePath, err) + } + if setupFn, ok := extensionRuntime.state.Get(-1).(*lua.LFunction); ok { + extensionRuntime.state.Push(setupFn) + extensionRuntime.state.Push(extensionRuntime.state.GetGlobal("librecode")) + if err := extensionRuntime.state.PCall(1, 0, nil); err != nil { + manager.unregisterRuntime(extensionRuntime) + extensionRuntime.state.Close() + return fmt.Errorf("extension: setup %s: %w", absolutePath, err) + } + } + recordLuaCallDuration(extensionRuntime, startedAt) + + manager.lock.Lock() + manager.extensions = append(manager.extensions, extensionRuntime) + manager.lock.Unlock() + manager.logger.Debug("loaded lua extension", slog.String("path", absolutePath)) + + return nil +} + +func (manager *Manager) installAPI(extensionRuntime *luaExtension) { + apiTable := extensionRuntime.state.NewTable() + extensionRuntime.state.SetFuncs(apiTable, map[string]lua.LGFunction{ + "register_command": manager.luaRegisterCommand(extensionRuntime), + "register_tool": manager.luaRegisterTool(extensionRuntime), + "on": manager.luaOn(extensionRuntime), + "log": manager.luaLog(extensionRuntime), + }) + extensionRuntime.state.SetField(apiTable, "api", manager.luaCoreAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "autocmd", manager.luaAutocmdAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "buf", manager.luaBufferAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "command", manager.luaCommandAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "event", manager.luaEventAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "action", manager.luaActionAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "timer", manager.luaTimerAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "keymap", manager.luaKeymapAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "layout", manager.luaLayoutAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "ui", manager.luaUIAPI(extensionRuntime)) + extensionRuntime.state.SetField(apiTable, "win", manager.luaWindowAPI(extensionRuntime)) + extensionRuntime.state.SetGlobal("librecode", apiTable) + extensionRuntime.state.PreloadModule("librecode", func(state *lua.LState) int { + state.Push(apiTable) + + return 1 + }) +} + +func (manager *Manager) luaRegisterCommand(extensionRuntime *luaExtension) lua.LGFunction { + return func(state *lua.LState) int { + name, description, function := luaRegistrationArgs(state) + definition := Command{Name: name, Description: description, Extension: extensionRuntime.name} + + manager.lock.Lock() + manager.commands[name] = luaCommand{extension: extensionRuntime, function: function, definition: definition} + extensionRuntime.commands = append(extensionRuntime.commands, name) + manager.lock.Unlock() + + return 0 + } +} + +func (manager *Manager) luaRegisterTool(extensionRuntime *luaExtension) lua.LGFunction { + return func(state *lua.LState) int { + name, description, function := luaRegistrationArgs(state) + definition := Tool{ + InputSchema: luaOptionalSchema(state, 4), + Name: name, + Description: description, + Extension: extensionRuntime.name, + } + + manager.lock.Lock() + manager.tools[name] = luaTool{extension: extensionRuntime, function: function, definition: definition} + extensionRuntime.tools = append(extensionRuntime.tools, name) + manager.lock.Unlock() + + return 0 + } +} + +func (manager *Manager) luaOn(extensionRuntime *luaExtension) lua.LGFunction { + return func(state *lua.LState) int { + eventName := state.CheckString(1) + priority, function := luaEventHandlerArgs(state) + manager.registerHandler(extensionRuntime, eventName, priority, function) + + return 0 + } +} + +func (manager *Manager) luaLog(extensionRuntime *luaExtension) lua.LGFunction { + return func(state *lua.LState) int { + message := state.CheckString(1) + manager.logger.Info( + "lua extension", + slog.String("extension", extensionRuntime.name), + slog.String("message", message), + ) + + return 0 + } +} + +func luaRegistrationArgs(state *lua.LState) (name, description string, function *lua.LFunction) { + return state.CheckString(1), state.OptString(2, ""), state.CheckFunction(3) +} + +func luaOptionalSchema(state *lua.LState, index int) map[string]any { + if state.GetTop() < index { + return map[string]any{} + } + if table, ok := state.Get(index).(*lua.LTable); ok { + return luaTableToMap(table) + } + + return map[string]any{} +} + +func luaEventHandlerArgs(state *lua.LState) (priority int, function *lua.LFunction) { + if handler, ok := state.Get(2).(*lua.LFunction); ok { + return 0, handler + } + + options := state.CheckTable(2) + + return int(lua.LVAsNumber(options.RawGetString("priority"))), state.CheckFunction(3) +} + +func openExtensionLibs(state *lua.LState) { + libraries := []struct { + open lua.LGFunction + name string + }{ + {name: lua.BaseLibName, open: lua.OpenBase}, + {name: lua.LoadLibName, open: lua.OpenPackage}, + {name: lua.TabLibName, open: lua.OpenTable}, + {name: lua.StringLibName, open: lua.OpenString}, + {name: lua.MathLibName, open: lua.OpenMath}, + {name: lua.IoLibName, open: lua.OpenIo}, + {name: lua.OsLibName, open: lua.OpenOs}, + {name: lua.DebugLibName, open: lua.OpenDebug}, + } + + for _, library := range libraries { + state.Push(state.NewFunction(library.open)) + state.Push(lua.LString(library.name)) + state.Call(1, 0) + } +} + +func extensionName(extensionPath string) string { + baseName := filepath.Base(extensionPath) + return strings.TrimSuffix(baseName, filepath.Ext(baseName)) +} diff --git a/internal/extension/manager_sources.go b/internal/extension/manager_sources.go new file mode 100644 index 00000000..2720ff21 --- /dev/null +++ b/internal/extension/manager_sources.go @@ -0,0 +1,156 @@ +package extension + +import ( + "fmt" + "os" + "path/filepath" + "sort" + "strings" + + lua "github.com/yuin/gopher-lua" +) + +type luaSource struct { + Path string + Manifest bool +} + +func discoverLuaSources(extensionPath string) ([]luaSource, error) { + if extensionPath == "" { + return []luaSource{}, nil + } + + info, err := os.Stat(extensionPath) + if err != nil { + if os.IsNotExist(err) { + return []luaSource{}, nil + } + return nil, fmt.Errorf("extension: stat %s: %w", extensionPath, err) + } + + if !info.IsDir() { + if strings.HasSuffix(extensionPath, ".lua") { + return []luaSource{{Path: extensionPath, Manifest: false}}, nil + } + return []luaSource{}, nil + } + + return discoverLuaDir(extensionPath) +} + +func discoverLuaDir(root string) ([]luaSource, error) { + manifestPath := filepath.Join(root, "init.lua") + if info, err := os.Stat(manifestPath); err == nil && !info.IsDir() { + return []luaSource{{Path: manifestPath, Manifest: true}}, nil + } + + sources := []luaSource{} + walkErr := filepath.WalkDir(root, func(currentPath string, dirEntry os.DirEntry, walkErr error) error { + return collectLuaSource(root, currentPath, dirEntry, walkErr, &sources) + }) + if walkErr != nil { + return nil, fmt.Errorf("extension: walk %s: %w", root, walkErr) + } + sort.Slice(sources, func(i, j int) bool { return sources[i].Path < sources[j].Path }) + + return sources, nil +} + +func collectLuaSource(root, currentPath string, dirEntry os.DirEntry, walkErr error, sources *[]luaSource) error { + if walkErr != nil { + return walkErr + } + if isExtensionDir(root, currentPath, dirEntry) { + return collectLuaSourceDir(root, currentPath, sources) + } + if strings.HasSuffix(currentPath, ".lua") { + *sources = append(*sources, luaSource{Path: currentPath, Manifest: false}) + } + + return nil +} + +func isExtensionDir(root, currentPath string, dirEntry os.DirEntry) bool { + if dirEntry.IsDir() { + return true + } + if currentPath == root || dirEntry.Type()&os.ModeSymlink == 0 { + return false + } + info, err := os.Stat(currentPath) + + return err == nil && info.IsDir() +} + +func collectLuaSourceDir(root, currentPath string, sources *[]luaSource) error { + if currentPath == root { + return nil + } + if filepath.Base(currentPath) == "lua" { + return filepath.SkipDir + } + manifestPath := filepath.Join(currentPath, "init.lua") + if info, err := os.Stat(manifestPath); err == nil && !info.IsDir() { + *sources = append(*sources, luaSource{Path: manifestPath, Manifest: true}) + return filepath.SkipDir + } + + return nil +} + +func (manager *Manager) addModuleRootsForPath(extensionPath string) { + if strings.TrimSpace(extensionPath) == "" { + return + } + absolutePath, err := filepath.Abs(extensionPath) + if err != nil { + absolutePath = extensionPath + } + roots := moduleRootsForPath(absolutePath) + + manager.lock.Lock() + defer manager.lock.Unlock() + + seen := make(map[string]struct{}, len(manager.moduleRoots)+len(roots)) + for _, root := range manager.moduleRoots { + seen[root] = struct{}{} + } + for _, root := range roots { + if _, ok := seen[root]; ok { + continue + } + manager.moduleRoots = append(manager.moduleRoots, root) + seen[root] = struct{}{} + } +} + +func moduleRootsForPath(extensionPath string) []string { + root := extensionPath + if info, err := os.Stat(extensionPath); err == nil && !info.IsDir() { + root = filepath.Dir(extensionPath) + } + + return []string{root} +} + +func (manager *Manager) configurePackagePath(state *lua.LState) { + packageTable, ok := state.GetGlobal("package").(*lua.LTable) + if !ok { + return + } + patterns := []string{packageTable.RawGetString("path").String()} + for _, root := range manager.moduleRootsSnapshot() { + patterns = append(patterns, + filepath.ToSlash(filepath.Join(root, "?.lua")), + filepath.ToSlash(filepath.Join(root, "?", "init.lua")), + ) + } + packageTable.RawSetString("path", lua.LString(strings.Join(patterns, ";"))) +} + +func (manager *Manager) moduleRootsSnapshot() []string { + manager.lock.RLock() + defer manager.lock.RUnlock() + + return append([]string{}, manager.moduleRoots...) +} diff --git a/internal/extension/manager_test.go b/internal/extension/manager_test.go index b53f2997..a213028b 100644 --- a/internal/extension/manager_test.go +++ b/internal/extension/manager_test.go @@ -8,6 +8,7 @@ import ( "path/filepath" "runtime" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -548,6 +549,92 @@ end`, assert.Equal(t, "linked", manager.Extensions()[0].Name) } +func TestManager_DoesNotRunCanceledCommandOrTool(t *testing.T) { + t.Parallel() + + manager := loadTestExtension(t, ` +local lc = require("librecode") + +lc.register_command("touch", "Touch command", function() + lc.buf.set_text("side_effect", "command") + return "command" +end) + +lc.register_tool("touch", "Touch tool", function() + lc.buf.set_text("side_effect", "tool") + return { content = "tool" } +end) +`) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, commandErr := manager.ExecuteCommand(ctx, "touch", "") + require.ErrorIs(t, commandErr, context.Canceled) + _, toolErr := manager.ExecuteTool(ctx, "touch", map[string]any{}) + require.ErrorIs(t, toolErr, context.Canceled) +} + +func TestManager_RollsBackPartialRegistrationsOnLoadFailure(t *testing.T) { + t.Parallel() + + manager := extension.NewManager(slog.New(slog.NewTextHandler(io.Discard, nil))) + t.Cleanup(manager.Shutdown) + extensionPath := filepath.Join(t.TempDir(), "broken.lua") + require.NoError(t, writeTestFile(extensionPath, ` +local lc = require("librecode") + +lc.register_command("broken", "Broken command", function() + return "should not run" +end) +lc.register_tool("broken", "Broken tool", function() + return { content = "should not run" } +end) +lc.on("startup", function() end) +lc.keymap.set({ focus = "composer" }, "x", function() return true end) +lc.timer.defer(1000, function() end) +error("boom") +`)) + + err := manager.LoadFile(context.Background(), extensionPath) + require.ErrorContains(t, err, "extension: load") + require.ErrorContains(t, err, "boom") + assert.Empty(t, manager.Commands()) + assert.Empty(t, manager.Tools()) + assert.Empty(t, manager.Extensions()) + assert.False(t, manager.HasTerminalEventHandlers(testEventStartup)) + assert.False(t, manager.HasTerminalEventHandlers(testEventKey)) + _, hasTimer := manager.NextTimerDelay(time.Now()) + assert.False(t, hasTimer) +} + +func TestManager_RollsBackPartialRegistrationsOnSetupFailure(t *testing.T) { + t.Parallel() + + extensionRoot := t.TempDir() + require.NoError(t, writeTestFile( + filepath.Join(extensionRoot, "init.lua"), + `return { name = "broken", entry = "main.lua" }`, + )) + require.NoError(t, writeTestFile(filepath.Join(extensionRoot, "main.lua"), ` +return function(librecode) + librecode.register_command("broken", "Broken command", function() + return "should not run" + end) + librecode.on("startup", function() end) + error("boom") +end +`)) + + manager := extension.NewManager(slog.New(slog.NewTextHandler(io.Discard, nil))) + t.Cleanup(manager.Shutdown) + err := manager.LoadPaths(context.Background(), []string{extensionRoot}) + require.ErrorContains(t, err, "extension: setup") + require.ErrorContains(t, err, "boom") + assert.Empty(t, manager.Commands()) + assert.Empty(t, manager.Extensions()) + assert.False(t, manager.HasTerminalEventHandlers(testEventStartup)) +} + func assertLoadedCommand(t *testing.T, commands []extension.Command, extensionName string) { t.Helper()