diff --git a/go.mod b/go.mod index d10aa59..32b63d9 100644 --- a/go.mod +++ b/go.mod @@ -5,7 +5,6 @@ go 1.24.2 require ( github.com/hashicorp/go-hclog v1.6.3 github.com/hashicorp/go-plugin v1.6.3 - github.com/magodo/slog2hclog v0.0.0-20240614031327-090ebd72a033 github.com/zeebo/errs/v2 v2.0.5 golang.org/x/sys v0.33.0 google.golang.org/grpc v1.72.2 diff --git a/go.sum b/go.sum index 2671b5d..d63083b 100644 --- a/go.sum +++ b/go.sum @@ -24,8 +24,6 @@ github.com/hashicorp/yamux v0.1.2 h1:XtB8kyFOyHXYVFnwT5C3+Bdo8gArse7j2AQ0DA0Uey8 github.com/hashicorp/yamux v0.1.2/go.mod h1:C+zze2n6e/7wshOZep2A70/aQU6QBRWJO/G6FT1wIns= github.com/jhump/protoreflect v1.15.1 h1:HUMERORf3I3ZdX05WaQ6MIpd/NJ434hTp5YiKgfCL6c= github.com/jhump/protoreflect v1.15.1/go.mod h1:jD/2GMKKE6OqX8qTjhADU1e6DShO+gavG9e0Q693nKo= -github.com/magodo/slog2hclog v0.0.0-20240614031327-090ebd72a033 h1:K2seYsMAzoICCLdDe7uU2WyaACLW+tvdTWG3QB+pyec= -github.com/magodo/slog2hclog v0.0.0-20240614031327-090ebd72a033/go.mod h1:8PvdX1kpjMEmR7LTNZ0QulFpD9j3E/eJijru+4nHY7M= github.com/mattn/go-colorable v0.1.9/go.mod h1:u6P/XSegPjTcexA+o6vUJrdnUu04hMope9wVRipJSqc= github.com/mattn/go-colorable v0.1.12/go.mod h1:u5H1YNBxpqRaxsYJYSkiCWKzEfiAb1Gb520KVy5xxl4= github.com/mattn/go-colorable v0.1.14 h1:9A9LHSqF/7dyVVX6g0U9cwm9pG3kP9gSzcuIPHPsaIE= diff --git a/internal/bootstrap/register.go b/internal/bootstrap/register.go index 0643f79..de99703 100644 --- a/internal/bootstrap/register.go +++ b/internal/bootstrap/register.go @@ -19,7 +19,7 @@ type HostDialer interface { // register given servers with the gRPC server. The given dialer and logger will // be used when the plugins are initialized. -func register(s *grpc.Server, servers []api.ServiceServer, logger hclog.Logger, dialer HostDialer) { +func Register(s *grpc.Server, servers []api.ServiceServer, logger hclog.Logger, dialer HostDialer) { var names []string var impls []any for _, server := range servers { diff --git a/internal/bootstrap/register_test.go b/internal/bootstrap/register_test.go index b8816b7..4e1ae56 100644 --- a/internal/bootstrap/register_test.go +++ b/internal/bootstrap/register_test.go @@ -183,7 +183,7 @@ func TestRegister(t *testing.T) { grpcSrv := grpc.NewServer() // Act - register(grpcSrv, tc.svcs, hclog.Default(), hostDialerMock{}) + Register(grpcSrv, tc.svcs, hclog.Default(), hostDialerMock{}) }) } } diff --git a/internal/bootstrap/serve.go b/internal/bootstrap/serve.go index 2bfe28f..29343fb 100644 --- a/internal/bootstrap/serve.go +++ b/internal/bootstrap/serve.go @@ -40,7 +40,7 @@ func newHCPlugin(logger hclog.Logger, pluginServer api.PluginServer, serviceServ } func (p *hcServer) GRPCServer(broker *goplugin.GRPCBroker, server *grpc.Server) (err error) { - register(server, p.servers, p.logger, &hcDialer{broker: broker}) + Register(server, p.servers, p.logger, &hcDialer{broker: broker}) return nil } diff --git a/internal/slog2hclog/logger_test.go b/internal/slog2hclog/logger_test.go new file mode 100644 index 0000000..263e675 --- /dev/null +++ b/internal/slog2hclog/logger_test.go @@ -0,0 +1,87 @@ +package slog2hclog + +import ( + "log/slog" + "testing" + + "github.com/hashicorp/go-hclog" +) + +func TestSetLogLevel(t *testing.T) { + // create test cases + tests := []struct { + name string + input string + want hclog.Level + }{ + { + name: "zero values", + want: hclog.Info, + }, { + name: "invalid input", + input: "invalid", + want: hclog.Info, + }, { + name: "debug lower case", + input: "debug", + want: hclog.Debug, + }, { + name: "debug upper case", + input: "DEBUG", + want: hclog.Debug, + }, { + name: "debug mixed case", + input: "DEbug", + want: hclog.Debug, + }, { + name: "info lower case", + input: "info", + want: hclog.Info, + }, { + name: "info upper case", + input: "INFO", + want: hclog.Info, + }, { + name: "info mixed case", + input: "INfo", + want: hclog.Info, + }, { + name: "warn lower case", + input: "warn", + want: hclog.Warn, + }, { + name: "warn upper case", + input: "WARN", + want: hclog.Warn, + }, { + name: "warn mixed case", + input: "WArn", + want: hclog.Warn, + }, { + name: "error lower case", + input: "error", + want: hclog.Error, + }, { + name: "error upper case", + input: "ERROR", + want: hclog.Error, + }, { + name: "error mixed case", + input: "ERRor", + want: hclog.Error, + }, + } + + // run the tests + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + // Act + log := NewWithLevel(slog.Default(), tc.input) + + // Assert + if got := log.GetLevel(); got != tc.want { + t.Errorf("expected value: %v, got: %s", tc.want, got) + } + }) + } +} diff --git a/internal/slog2hclog/slog2hclog.go b/internal/slog2hclog/slog2hclog.go new file mode 100644 index 0000000..316a82f --- /dev/null +++ b/internal/slog2hclog/slog2hclog.go @@ -0,0 +1,245 @@ +package slog2hclog + +import ( + "context" + "io" + "log" + "log/slog" + "sort" + "strings" + + "github.com/hashicorp/go-hclog" +) + +type stdslogWrapper struct { + slog *slog.Logger + oriSlog *slog.Logger + lvar *slog.LevelVar + names []string + args []interface{} +} + +func (s *stdslogWrapper) clone() *stdslogWrapper { + newSlog := *s.slog + var oriSlog *slog.Logger = nil + if s.oriSlog != nil { + newOriSlog := *s.oriSlog + oriSlog = &newOriSlog + } + return &stdslogWrapper{ + slog: &newSlog, + oriSlog: oriSlog, + names: append([]string{}, s.names...), + args: append([]interface{}{}, s.args...), + } +} + +var _ hclog.Logger = &stdslogWrapper{} + +const ( + SlogLevelTrace = slog.LevelDebug - 4 + SlogLevelOff = slog.LevelError + 4 +) + +var levelMapToSlog = map[hclog.Level]slog.Level{ + hclog.Off: SlogLevelOff, + hclog.Error: slog.LevelError, + hclog.Warn: slog.LevelWarn, + hclog.Info: slog.LevelInfo, + hclog.Debug: slog.LevelDebug, + hclog.Trace: SlogLevelTrace, +} + +var levelMapFromSlog = map[slog.Level]hclog.Level{ + SlogLevelOff: hclog.Off, + slog.LevelError: hclog.Error, + slog.LevelWarn: hclog.Warn, + slog.LevelInfo: hclog.Info, + slog.LevelDebug: hclog.Debug, + SlogLevelTrace: hclog.Trace, +} + +type levelOverrideHandler struct { + slog.Handler + overrideLevel slog.Level +} + +func (h *levelOverrideHandler) Enabled(_ context.Context, level slog.Level) bool { + // Force-enable only messages equal or above overrideLevel + return level >= h.overrideLevel +} + +// New wraps a slog.Logger to a hclog.Logger with info log level. +func New(l *slog.Logger) hclog.Logger { + return NewWithLevel(l, "info") +} + +// NewWithLevel wraps a slog.Logger to a hclog.Logger. +func NewWithLevel(l *slog.Logger, logLevel string) hclog.Logger { + level := hclog.LevelFromString(logLevel) + wrapper := &stdslogWrapper{ + slog: slog.New(&levelOverrideHandler{ + Handler: l.Handler(), + overrideLevel: levelMapToSlog[level], + }), + oriSlog: nil, + lvar: new(slog.LevelVar), + names: []string{}, + args: []interface{}{}, + } + wrapper.SetLevel(level) + + return wrapper +} + +func (s *stdslogWrapper) Trace(msg string, args ...interface{}) { + s.slog.Log(context.Background(), SlogLevelTrace, msg, args...) +} + +func (s *stdslogWrapper) Debug(msg string, args ...interface{}) { + s.slog.Debug(msg, args...) +} + +func (s *stdslogWrapper) Info(msg string, args ...interface{}) { + s.slog.Info(msg, args...) +} + +func (s *stdslogWrapper) Warn(msg string, args ...interface{}) { + s.slog.Warn(msg, args...) +} + +func (s *stdslogWrapper) Error(msg string, args ...interface{}) { + s.slog.Error(msg, args...) +} + +func (s *stdslogWrapper) GetLevel() hclog.Level { + if s.lvar == nil { + // lvar not set indicates the source slog.Logger has a fixed log level (or a default level, which equals to Info). + // In this case, we enumerate the log levels from lowest (Trace) to get the effective log level. + return s.getLowestLevel() + } + return levelMapFromSlog[s.lvar.Level()] +} + +// SetLevel only applies when the source slog.Logger has a slog.LevelVar level. +func (s *stdslogWrapper) SetLevel(level hclog.Level) { + if s.lvar != nil { + s.lvar.Set(levelMapToSlog[level]) + } +} + +func (s *stdslogWrapper) IsTrace() bool { + return s.slog.Enabled(context.Background(), SlogLevelTrace) +} + +func (s *stdslogWrapper) IsDebug() bool { + return s.slog.Enabled(context.Background(), slog.LevelDebug) +} + +func (s *stdslogWrapper) IsInfo() bool { + return s.slog.Enabled(context.Background(), slog.LevelInfo) +} + +func (s *stdslogWrapper) IsWarn() bool { + return s.slog.Enabled(context.Background(), slog.LevelWarn) +} + +func (s *stdslogWrapper) IsError() bool { + return s.slog.Enabled(context.Background(), slog.LevelError) +} + +func (s *stdslogWrapper) Log(level hclog.Level, msg string, args ...interface{}) { + s.slog.Log(context.Background(), levelMapToSlog[level], msg, args...) +} + +func (s *stdslogWrapper) Name() string { + return strings.Join(s.names, ".") +} + +func (s *stdslogWrapper) Named(name string) hclog.Logger { + sl := s.clone() + if len(s.names) == 0 { + newSlog := *sl.slog + sl.oriSlog = &newSlog + } + sl.names = append(sl.names, name) + sl.slog = slog.New(&levelOverrideHandler{ + Handler: s.slog.WithGroup(name).Handler(), + overrideLevel: levelMapToSlog[s.GetLevel()], + }) + return sl +} + +func (s *stdslogWrapper) ResetNamed(name string) hclog.Logger { + sl := s.clone() + + // Empty name indicates to clear the name + if name == "" { + if len(sl.names) == 0 { + return sl + } + sl.names = []string{} + sl.slog = sl.oriSlog + sl.oriSlog = nil + return sl + } + + // Non-empty name indicates to set the name + if len(sl.names) == 0 { + return sl.Named(name) + } + sl.names = []string{} + sl.slog = sl.oriSlog + sl.oriSlog = nil + return sl.Named(name) +} + +func (s *stdslogWrapper) With(args ...interface{}) hclog.Logger { + sl := s.clone() + sl.slog = s.slog.With(args...) + sl.args = append(sl.args, args...) + return sl +} + +func (s *stdslogWrapper) ImpliedArgs() []interface{} { + return s.args +} + +func (s *stdslogWrapper) StandardLogger(opts *hclog.StandardLoggerOptions) *log.Logger { + if opts == nil { + opts = &hclog.StandardLoggerOptions{} + } + + return log.New(s.StandardWriter(opts), "", 0) +} + +func (s *stdslogWrapper) StandardWriter(opts *hclog.StandardLoggerOptions) io.Writer { + newLog := s.clone() + return &stdlogAdapter{ + log: newLog, + inferLevels: opts.InferLevels, + inferLevelsWithTimestamp: opts.InferLevelsWithTimestamp, + forceLevel: opts.ForceLevel, + } +} + +func (s *stdslogWrapper) getLowestLevel() hclog.Level { + ctx := context.Background() + + var slogLvls []slog.Level + for lvlSlog := range levelMapFromSlog { + slogLvls = append(slogLvls, lvlSlog) + } + // Sort the slog levels from Trace up to Error + sort.Slice(slogLvls, func(i, j int) bool { + return int(slogLvls[i]) < int(slogLvls[j]) + }) + + for _, lvlSlog := range slogLvls { + lvl := levelMapFromSlog[lvlSlog] + if s.slog.Enabled(ctx, lvlSlog) { + return lvl + } + } + return hclog.Off +} diff --git a/internal/slog2hclog/stdlog.go b/internal/slog2hclog/stdlog.go new file mode 100644 index 0000000..7546027 --- /dev/null +++ b/internal/slog2hclog/stdlog.go @@ -0,0 +1,91 @@ +package slog2hclog + +import ( + "bytes" + "regexp" + "strings" + + "github.com/hashicorp/go-hclog" +) + +// Regex to ignore characters commonly found in timestamp formats from the +// beginning of inputs. +var logTimestampRegexp = regexp.MustCompile(`^[\d\s\:\/\.\+-TZ]*`) + +// Provides a io.Writer to shim the data out of *log.Logger +// and back into our Logger. This is basically the only way to +// build upon *log.Logger. +type stdlogAdapter struct { + log hclog.Logger + inferLevels bool + inferLevelsWithTimestamp bool + forceLevel hclog.Level +} + +// Take the data, infer the levels if configured, and send it through +// a regular Logger. +func (s *stdlogAdapter) Write(data []byte) (int, error) { + str := string(bytes.TrimRight(data, " \t\n")) + + if s.forceLevel != hclog.NoLevel { + // Use pickLevel to strip log levels included in the line since we are + // forcing the level + _, str := s.pickLevel(str) + + // Log at the forced level + s.dispatch(str, s.forceLevel) + } else if s.inferLevels { + if s.inferLevelsWithTimestamp { + str = s.trimTimestamp(str) + } + + level, str := s.pickLevel(str) + s.dispatch(str, level) + } else { + s.log.Info(str) + } + + return len(data), nil +} + +func (s *stdlogAdapter) dispatch(str string, level hclog.Level) { + switch level { + case hclog.Trace: + s.log.Trace(str) + case hclog.Debug: + s.log.Debug(str) + case hclog.Info: + s.log.Info(str) + case hclog.Warn: + s.log.Warn(str) + case hclog.Error: + s.log.Error(str) + default: + s.log.Info(str) + } +} + +// Detect, based on conventions, what log level this is. +func (s *stdlogAdapter) pickLevel(str string) (hclog.Level, string) { + switch { + case strings.HasPrefix(str, "[DEBUG]"): + return hclog.Debug, strings.TrimSpace(str[7:]) + case strings.HasPrefix(str, "[TRACE]"): + return hclog.Trace, strings.TrimSpace(str[7:]) + case strings.HasPrefix(str, "[INFO]"): + return hclog.Info, strings.TrimSpace(str[6:]) + case strings.HasPrefix(str, "[WARN]"): + return hclog.Warn, strings.TrimSpace(str[6:]) + case strings.HasPrefix(str, "[ERROR]"): + return hclog.Error, strings.TrimSpace(str[7:]) + case strings.HasPrefix(str, "[ERR]"): + return hclog.Error, strings.TrimSpace(str[5:]) + default: + return hclog.Info, str + } +} + +func (s *stdlogAdapter) trimTimestamp(str string) string { + idx := logTimestampRegexp.FindStringIndex(str) + return str[idx[1]:] +} diff --git a/pkg/catalog/builtin.go b/pkg/catalog/builtin.go new file mode 100644 index 0000000..2407a48 --- /dev/null +++ b/pkg/catalog/builtin.go @@ -0,0 +1,185 @@ +package catalog + +import ( + "context" + "errors" + "io" + "log/slog" + "sync" + "time" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + + "github.com/openkcm/plugin-sdk/api" + "github.com/openkcm/plugin-sdk/internal/bootstrap" + "github.com/openkcm/plugin-sdk/internal/slog2hclog" +) + +type BuiltIn struct { + Name string + Tags []string + Plugin api.PluginServer + Services []api.ServiceServer +} + +func MakeBuiltIn(name string, pluginServer api.PluginServer, serviceServers ...api.ServiceServer) BuiltIn { + return BuiltIn{ + Name: name, + Plugin: pluginServer, + Services: serviceServers, + } +} + +func loadBuiltIn(ctx context.Context, builtIn BuiltIn, pluginConfig PluginConfig) (_ *Plugin, err error) { + dialer := &builtinDialer{ + pluginName: builtIn.Name, + log: pluginConfig.Logger, + hostServices: pluginConfig.HostServices, + } + + var closers closerGroup + defer func() { + if err != nil { + _ = closers.Close() + } + }() + closers = append(closers, dialer) + + builtinServer, serverCloser := newBuiltInServer(pluginConfig.Logger) + closers = append(closers, serverCloser) + + pluginServers := append([]api.ServiceServer{builtIn.Plugin}, builtIn.Services...) + + log := slog2hclog.NewWithLevel(pluginConfig.Logger, pluginConfig.LogLevel) + bootstrap.Register(builtinServer, pluginServers, log, dialer) + + builtinConn, err := startPipeServer(builtinServer, pluginConfig.Logger) + if err != nil { + return nil, err + } + closers = append(closers, builtinConn) + + info := pluginInfo{ + name: builtIn.Name, + typ: builtIn.Plugin.Type(), + tags: builtIn.Tags, + } + + return newPlugin(ctx, builtinConn, info, pluginConfig.Logger, closers, pluginConfig.HostServices) +} + +func newBuiltInServer(log *slog.Logger) (*grpc.Server, io.Closer) { + drain := &drainHandlers{} + return grpc.NewServer( + grpc.ChainStreamInterceptor(drain.StreamServerInterceptor, streamPanicInterceptor(log)), + grpc.ChainUnaryInterceptor(drain.UnaryServerInterceptor, unaryPanicInterceptor(log)), + ), closerFunc(drain.Wait) +} + +type builtinDialer struct { + pluginName string + log *slog.Logger + hostServices []api.ServiceServer + conn *pipeConn +} + +func (d *builtinDialer) DialHost(context.Context) (grpc.ClientConnInterface, error) { + if d.conn != nil { + return d.conn, nil + } + server := newHostServer(d.log, d.pluginName) + conn, err := startPipeServer(server, d.log) + if err != nil { + return nil, err + } + d.conn = conn + return d.conn, nil +} + +func (d *builtinDialer) Close() error { + if d.conn != nil { + return d.conn.Close() + } + return nil +} + +type pipeConn struct { + grpc.ClientConnInterface + io.Closer +} + +func startPipeServer(server *grpc.Server, log *slog.Logger) (*pipeConn, error) { + pipeNet := newPipeNet() + + var wg sync.WaitGroup + + var closers closerGroup + closers = append(closers, closerFunc(wg.Wait), closerFunc(func() { + if !gracefulStopWithTimeout(server, time.Minute) { + log.Warn("Forced timed-out plugin server to stop") + } + }), closerFunc(func() { + err := pipeNet.Close() + if err != nil { + return + } + })) + + wg.Add(1) + go func() { + defer wg.Done() + if err := server.Serve(pipeNet); err != nil && !errors.Is(err, grpc.ErrServerStopped) { + log.Error("Pipe server unexpectedly failed to serve", "error", err) + } + }() + + // Dial the server + conn, err := grpc.NewClient( + "passthrough:IGNORED", + grpc.WithTransportCredentials(insecure.NewCredentials()), + grpc.WithContextDialer(pipeNet.DialContext), + ) + if err != nil { + return nil, err + } + closers = append(closers, conn) + + return &pipeConn{ + ClientConnInterface: conn, + Closer: closers, + }, nil +} + +type drainHandlers struct { + wg sync.WaitGroup +} + +func (d *drainHandlers) Wait() { + done := make(chan struct{}) + + go func() { + d.wg.Wait() + close(done) + }() + + t := time.NewTimer(time.Minute) + defer t.Stop() + + select { + case <-done: + case <-t.C: + } +} + +func (d *drainHandlers) UnaryServerInterceptor(ctx context.Context, req any, _ *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + d.wg.Add(1) + defer d.wg.Done() + return handler(ctx, req) +} + +func (d *drainHandlers) StreamServerInterceptor(srv any, ss grpc.ServerStream, _ *grpc.StreamServerInfo, handler grpc.StreamHandler) error { + d.wg.Add(1) + defer d.wg.Done() + return handler(srv, ss) +} diff --git a/pkg/catalog/catalog.go b/pkg/catalog/catalog.go index d9f4301..c9d74ce 100644 --- a/pkg/catalog/catalog.go +++ b/pkg/catalog/catalog.go @@ -8,7 +8,6 @@ import ( "time" "github.com/openkcm/plugin-sdk/api" - "github.com/openkcm/plugin-sdk/pkg/telemetry" ) const ( @@ -57,7 +56,7 @@ func (c *Catalog) LookupByTypeAndName(pluginType, pluginName string) *Plugin { return nil } -func Load(ctx context.Context, config Config) (catalog *Catalog, err error) { +func Load(ctx context.Context, config Config, builtIns ...BuiltIn) (catalog *Catalog, err error) { closers := make(closerGroup, 0) defer func() { // If loading fails, clear out the catalog and close down all plugins @@ -69,39 +68,90 @@ func Load(ctx context.Context, config Config) (catalog *Catalog, err error) { } }() + // in case if configuration logger is not set get the default one + if config.Logger == nil { + config.Logger = slog.Default() + } + configurers := make(Configurers, 0) for _, pluginConfig := range config.PluginConfigs { - pluginLog := config.Logger.With( - telemetry.PluginName, pluginConfig.Name, - telemetry.PluginType, pluginConfig.Type, - ) - pluginConfig.Logger = pluginLog - if pluginConfig.Disabled { config.Logger.Debug("Not loading plugin; disabled") continue } pluginConfig.HostServices = config.HostServices - plugin, err := loadPlugin(ctx, config.Logger, pluginConfig) + + plugin, err := loadPluginAs(ctx, config.Logger, pluginConfig, builtIns...) if err != nil { - config.Logger.ErrorContext(ctx, "Failed to load plugin", telemetry.PluginName, pluginConfig.Name, "error", err) - return nil, fmt.Errorf("failed to load plugin %q: %w", pluginConfig.Name, err) + return nil, err } - closers = append(closers, pluginCloser{plugin: plugin, log: pluginLog}) + + closers = append(closers, pluginCloser{plugin: plugin, log: plugin.Logger()}) cfgurer := makeConfigurer(plugin, pluginConfig) err = cfgurer.Configure(ctx) if err != nil { - config.Logger.ErrorContext(ctx, "Failed to configure plugin", telemetry.PluginName, pluginConfig.Name, "error", err) - return nil, fmt.Errorf("failed to configure plugin %q of type %q; %v", pluginConfig.Name, pluginConfig.Type, err) + return nil, fmt.Errorf("failed to configure plugin %s of type %s; %v", pluginConfig.Name, pluginConfig.Type, err) } configurers = append(configurers, cfgurer) - pluginLog.Info("Plugin loaded") + plugin.Logger().Info("Loaded plugin") } + return &Catalog{ closers: closers, configurers: configurers, }, nil } + +func loadPluginAs(ctx context.Context, logger *slog.Logger, pluginConfig PluginConfig, builtIns ...BuiltIn) (*Plugin, error) { + if pluginConfig.IsExternal() { + plugin, err := loadPluginAsExternal(ctx, logger, pluginConfig) + if err != nil { + return nil, fmt.Errorf("failed to load external plugin %s: %w", pluginConfig.Name, err) + } + return plugin, nil + } + + plugin, err := loadPluginAsBuiltIn(ctx, logger, pluginConfig, builtIns...) + if err != nil { + return nil, fmt.Errorf("failed to load builtin plugin %s: %w", pluginConfig.Name, err) + } + return plugin, nil +} + +func loadPluginAsExternal(ctx context.Context, logger *slog.Logger, pluginConfig PluginConfig) (*Plugin, error) { + if pluginConfig.Name == "" { + return nil, fmt.Errorf("failed to load external plugin, missing name") + } + + if pluginConfig.Type == "" { + return nil, fmt.Errorf("failed to load external plugin %s, missing type", pluginConfig.Name) + } + + pluginConfig.Logger = logger.With( + Name, pluginConfig.Name, + Type, pluginConfig.Type, + ) + + return loadPlugin(ctx, pluginConfig) +} + +func loadPluginAsBuiltIn(ctx context.Context, logger *slog.Logger, pluginConfig PluginConfig, builtIns ...BuiltIn) (*Plugin, error) { + if pluginConfig.Name == "" { + return nil, fmt.Errorf("failed to load builtin plugin, missing name") + } + + for _, builtin := range builtIns { + if builtin.Name == pluginConfig.Name { + pluginConfig.Logger = logger.With( + Name, pluginConfig.Name, + Type, builtin.Plugin.Type(), + ) + + return loadBuiltIn(ctx, builtin, pluginConfig) + } + } + return nil, fmt.Errorf("builtin plugin %q not found", pluginConfig.Name) +} diff --git a/pkg/catalog/catalog_test.go b/pkg/catalog/catalog_test.go index e879a52..4e11751 100644 --- a/pkg/catalog/catalog_test.go +++ b/pkg/catalog/catalog_test.go @@ -28,6 +28,34 @@ func TestLoad(t *testing.T) { HostServices: nil, }, wantError: false, + }, { + name: "missing name", + config: Config{ + Logger: slog.Default(), + PluginConfigs: []PluginConfig{ + { + Path: "/does/not/exist1", + Type: "TestService", + Logger: slog.Default(), + }, + }, + HostServices: nil, + }, + wantError: true, + }, { + name: "missing type", + config: Config{ + Logger: slog.Default(), + PluginConfigs: []PluginConfig{ + { + Path: "/does/not/exist2", + Name: "TestService", + Logger: slog.Default(), + }, + }, + HostServices: nil, + }, + wantError: true, }, { name: "invalid path", config: Config{ @@ -50,6 +78,7 @@ func TestLoad(t *testing.T) { // testpluginbinary is built in the TestMain function Path: "./testpluginbinary", Type: "TestService", + Name: "TestService", Logger: slog.Default(), }, }, @@ -66,6 +95,7 @@ func TestLoad(t *testing.T) { // testpluginbinary is built in the TestMain function Path: "./testpluginbinary", Type: "TestService", + Name: "TestService", Logger: slog.Default(), Tags: []string{"feature1", "feature2"}, }, @@ -110,6 +140,7 @@ func TestLookupByType(t *testing.T) { // testpluginbinary is built in the TestMain function Path: "./testpluginbinary", Type: "TestService", + Name: "TestService", Logger: slog.Default(), }, }, diff --git a/pkg/catalog/names.go b/pkg/catalog/names.go new file mode 100644 index 0000000..37433fb --- /dev/null +++ b/pkg/catalog/names.go @@ -0,0 +1,6 @@ +package catalog + +const ( + Name = "pluginName" + Type = "pluginType" +) diff --git a/pkg/catalog/pipenet.go b/pkg/catalog/pipenet.go new file mode 100644 index 0000000..2944461 --- /dev/null +++ b/pkg/catalog/pipenet.go @@ -0,0 +1,66 @@ +package catalog + +import ( + "context" + "errors" + "net" + "sync" +) + +type pipeAddr struct{} + +func (pipeAddr) Network() string { return "pipe" } +func (pipeAddr) String() string { return "pipe" } + +type pipeNet struct { + accept chan net.Conn + closed chan struct{} + closeOnce sync.Once +} + +func newPipeNet() *pipeNet { + return &pipeNet{ + accept: make(chan net.Conn), + closed: make(chan struct{}), + } +} + +func (n *pipeNet) Addr() net.Addr { + return pipeAddr{} +} + +func (n *pipeNet) Accept() (net.Conn, error) { + select { + case s := <-n.accept: + return s, nil + case <-n.closed: + return nil, errors.New("closed") + } +} + +func (n *pipeNet) DialContext(ctx context.Context, _ string) (conn net.Conn, err error) { + c, s := net.Pipe() + + defer func() { + if err != nil { + _ = c.Close() + _ = s.Close() + } + }() + + select { + case <-ctx.Done(): + return nil, ctx.Err() + case n.accept <- s: + return c, nil + case <-n.closed: + return nil, errors.New("network closed") + } +} + +func (n *pipeNet) Close() error { + n.closeOnce.Do(func() { + close(n.closed) + }) + return nil +} diff --git a/pkg/catalog/plugin.go b/pkg/catalog/plugin.go index 6a81d79..4daca38 100644 --- a/pkg/catalog/plugin.go +++ b/pkg/catalog/plugin.go @@ -9,15 +9,14 @@ import ( "io" "log/slog" "os/exec" - "strings" - "github.com/magodo/slog2hclog" "google.golang.org/grpc" goplugin "github.com/hashicorp/go-plugin" "github.com/openkcm/plugin-sdk/api" "github.com/openkcm/plugin-sdk/internal/bootstrap" + "github.com/openkcm/plugin-sdk/internal/slog2hclog" ) type PluginConfigs []PluginConfig @@ -55,6 +54,14 @@ type PluginConfig struct { Tags []string } +func (c *PluginConfig) IsExternal() bool { + return c.Path != "" +} + +func (c PluginConfig) IsEnabled() bool { + return !c.Disabled +} + // PluginInfo provides the information for the loaded plugin. type PluginInfo interface { // The name of the plugin @@ -82,19 +89,17 @@ func (p *Plugin) ClientConnection() grpc.ClientConnInterface { func (p *Plugin) Info() PluginInfo { return p.info } - +func (p *Plugin) Logger() *slog.Logger { + return p.logger +} func (p *Plugin) GrpcServiceNames() []string { return p.grpcServiceNames } -func loadPlugin(ctx context.Context, logger *slog.Logger, config PluginConfig) (*Plugin, error) { - logger.InfoContext(ctx, "Loading plugin", "name", config.Name, "path", config.Path) - - logLevelPlugin := new(slog.LevelVar) - setLogLevel(logLevelPlugin, config.LogLevel) +func loadPlugin(ctx context.Context, config PluginConfig) (*Plugin, error) { + config.Logger.InfoContext(ctx, "Loading plugin", "name", config.Name, "path", config.Path) cmd := pluginCmd(config.Path, config.Args...) - injectEnv(config, cmd) // Create the secure config based on the (optional) checksum @@ -104,9 +109,10 @@ func loadPlugin(ctx context.Context, logger *slog.Logger, config PluginConfig) ( } // Start the plugin client + pluginClient := goplugin.NewClient(&goplugin.ClientConfig{ SecureConfig: seccfg, - Logger: slog2hclog.New(config.Logger, logLevelPlugin), + Logger: slog2hclog.NewWithLevel(config.Logger, config.LogLevel), HandshakeConfig: goplugin.HandshakeConfig{ ProtocolVersion: 1, MagicCookieKey: config.Type, @@ -182,12 +188,12 @@ type pluginCloser struct { } func (c pluginCloser) Close() error { - c.log.Debug("Unloading plugin") + c.log.Info("Plugins unloading") if err := c.plugin.Close(); err != nil { c.log.Error("Failed to unload plugin", "error", err) return err } - c.log.Info("Plugin unloaded") + c.log.Info("Plugins unloaded") return nil } @@ -248,20 +254,3 @@ func initPlugin(ctx context.Context, conn grpc.ClientConnInterface, hostServices defer cancel() return bootstrap.Init(ctx, conn, hostServiceGRPCServiceNames) } - -// setLogLevel converts the level string used in the config to a slog.LevelVar -// and sets the levelVar to the corresponding level. -func setLogLevel(levelVar *slog.LevelVar, level string) { - switch strings.ToLower(level) { - case "debug": - levelVar.Set(slog.LevelDebug) - case "info": - levelVar.Set(slog.LevelInfo) - case "warn": - levelVar.Set(slog.LevelWarn) - case "error": - levelVar.Set(slog.LevelError) - default: - levelVar.Set(slog.LevelInfo) - } -} diff --git a/pkg/catalog/plugin_test.go b/pkg/catalog/plugin_test.go index 78954d1..3027dea 100644 --- a/pkg/catalog/plugin_test.go +++ b/pkg/catalog/plugin_test.go @@ -1,7 +1,6 @@ package catalog import ( - "log/slog" "os/exec" "testing" ) @@ -127,87 +126,6 @@ func TestBuildSecureConfig(t *testing.T) { } } -func TestSetLogLevel(t *testing.T) { - lv := new(slog.LevelVar) - - // create test cases - tests := []struct { - name string - input string - want slog.Level - }{ - { - name: "zero values", - want: slog.LevelInfo, - }, { - name: "invalid input", - input: "invalid", - want: slog.LevelInfo, - }, { - name: "debug lower case", - input: "debug", - want: slog.LevelDebug, - }, { - name: "debug upper case", - input: "DEBUG", - want: slog.LevelDebug, - }, { - name: "debug mixed case", - input: "DEbug", - want: slog.LevelDebug, - }, { - name: "info lower case", - input: "info", - want: slog.LevelInfo, - }, { - name: "info upper case", - input: "INFO", - want: slog.LevelInfo, - }, { - name: "info mixed case", - input: "INfo", - want: slog.LevelInfo, - }, { - name: "warn lower case", - input: "warn", - want: slog.LevelWarn, - }, { - name: "warn upper case", - input: "WARN", - want: slog.LevelWarn, - }, { - name: "warn mixed case", - input: "WArn", - want: slog.LevelWarn, - }, { - name: "error lower case", - input: "error", - want: slog.LevelError, - }, { - name: "error upper case", - input: "ERROR", - want: slog.LevelError, - }, { - name: "error mixed case", - input: "ERRor", - want: slog.LevelError, - }, - } - - // run the tests - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - // Act - setLogLevel(lv, tc.input) - - // Assert - if got := lv.Level(); got != tc.want { - t.Errorf("expected value: %v, got: %s", tc.want, got) - } - }) - } -} - func TestInjectEnv(t *testing.T) { // Arrange cmd := &exec.Cmd{} diff --git a/pkg/telemetry/names.go b/pkg/telemetry/names.go deleted file mode 100644 index bbb036c..0000000 --- a/pkg/telemetry/names.go +++ /dev/null @@ -1,6 +0,0 @@ -package telemetry - -const ( - PluginName = "plugin_name" - PluginType = "plugin_type" -)