refactor: code cleanup, stricter config validation (#155)

This commit is contained in:
Henrique Dias
2024-07-21 20:52:50 +02:00
committed by GitHub
parent 46d54e4465
commit c125bedae1
19 changed files with 683 additions and 654 deletions
+84 -82
View File
@@ -1,7 +1,8 @@
package cmd
import (
"log"
"errors"
"fmt"
"net"
"net/http"
"os"
@@ -9,27 +10,21 @@ import (
"strings"
"syscall"
"github.com/hacdias/webdav/v4/lib"
"github.com/spf13/cobra"
v "github.com/spf13/viper"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
)
var (
cfgFile string
)
func init() {
cobra.OnInitialize(initConfig)
flags := rootCmd.Flags()
flags.StringVarP(&cfgFile, "config", "c", "", "config file path")
flags.BoolP("tls", "t", false, "enable tls")
flags.Bool("auth", true, "enable auth")
flags.String("cert", "cert.pem", "TLS certificate")
flags.String("key", "key.pem", "TLS key")
flags.StringP("address", "a", "0.0.0.0", "address to listen to")
flags.StringP("port", "p", "0", "port to listen to")
flags.StringP("config", "c", "", "config file path")
flags.BoolP("tls", "t", false, "enable TLS")
flags.Bool("auth", false, "enable authentication")
flags.String("cert", "cert.pem", "path to TLS certificate")
flags.String("key", "key.pem", "path to TLS key")
flags.StringP("address", "a", "0.0.0.0", "address to listen on")
flags.StringP("port", "p", "0", "port to listen on")
flags.StringP("prefix", "P", "/", "URL path prefix")
flags.String("log_format", "console", "logging format")
}
@@ -53,91 +48,98 @@ The precedence of the configuration values are as follows:
The environment variables are prefixed by "WD_" followed by the option
name in caps. So to set "cert" via an env variable, you should
set WD_CERT.`,
Run: func(cmd *cobra.Command, args []string) {
RunE: func(cmd *cobra.Command, args []string) error {
flags := cmd.Flags()
cfg := readConfig(flags)
// Build address and listener
laddr := getOpt(flags, "address")
var lnet string
if strings.HasPrefix(laddr, "unix:") {
laddr = laddr[5:]
lnet = "unix"
} else {
laddr = laddr + ":" + getOpt(flags, "port")
lnet = "tcp"
}
listener, err := net.Listen(lnet, laddr)
sigc := make(chan os.Signal, 1)
signal.Notify(sigc, os.Interrupt, syscall.SIGTERM)
go func(c chan os.Signal) {
// Wait for a SIGINT or SIGKILL:
sig := <-c
log.Printf("Caught signal %s: shutting down.", sig)
// Stop listening (and unlink the socket if unix type):
listener.Close()
// And we're done:
os.Exit(0)
}(sigc)
cfgFilename, _ := flags.GetString("config")
cfg, err := lib.ParseConfig(cfgFilename, flags)
if err != nil {
log.Fatal(err)
return err
}
loggerConfig := zap.NewProductionConfig()
loggerConfig.DisableCaller = true
if cfg.Debug {
loggerConfig.Level = zap.NewAtomicLevelAt(zap.DebugLevel)
}
loggerConfig.EncoderConfig.EncodeTime = zapcore.ISO8601TimeEncoder
loggerConfig.Encoding = cfg.LogFormat
logger, err := loggerConfig.Build()
// Create HTTP handler from the config
handler, err := lib.NewHandler(cfg)
if err != nil {
// if we fail to configure proper logging, then the user has deliberately
// misconfigured the logger. Abort.
panic(err)
return err
}
zap.ReplaceGlobals(logger)
// Setup the logger based on the configuration
err = setupLogger(cfg)
if err != nil {
return err
}
defer func() {
// Flush the logger at the end
_ = zap.L().Sync()
}()
// Tell the user the port in which is listening.
zap.L().Info("Listening", zap.String("address", listener.Addr().String()))
// Starts the server.
if getOptB(flags, "tls") {
if err := http.ServeTLS(listener, cfg, getOpt(flags, "cert"), getOpt(flags, "key")); err != nil {
zap.L().Fatal("shutting server", zap.Error(err))
}
} else {
if err := http.Serve(listener, cfg); err != nil {
zap.L().Fatal("shutting server", zap.Error(err))
}
// Build listener
listener, err := getListener(cfg)
if err != nil {
return err
}
// Trap exiting signals
quit := make(chan os.Signal, 1)
go func() {
zap.L().Info("listening", zap.String("address", listener.Addr().String()))
var err error
if cfg.TLS {
err = http.ServeTLS(listener, handler, cfg.Cert, cfg.Key)
} else {
err = http.Serve(listener, handler)
}
if err != nil && !errors.Is(err, http.ErrServerClosed) {
zap.L().Error("failed to start server", zap.Error(err))
}
quit <- os.Interrupt
}()
signal.Notify(quit, os.Interrupt, syscall.SIGTERM)
signal := <-quit
zap.L().Info("caught signal, shutting down", zap.Stringer("signal", signal))
_ = listener.Close()
return nil
},
}
func initConfig() {
if cfgFile == "" {
v.AddConfigPath(".")
v.AddConfigPath("/etc/webdav/")
v.SetConfigName("config")
func getListener(cfg *lib.Config) (net.Listener, error) {
var (
address string
network string
)
if strings.HasPrefix(cfg.Address, "unix:") {
address = cfg.Address[5:]
network = "unix"
} else {
v.SetConfigFile(cfgFile)
address = fmt.Sprintf("%s:%d", cfg.Address, cfg.Port)
network = "tcp"
}
v.SetEnvPrefix("WD")
v.AutomaticEnv()
v.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
return net.Listen(network, address)
}
if err := v.ReadInConfig(); err != nil {
if _, ok := err.(v.ConfigParseError); ok {
panic(err)
}
cfgFile = "No config file used"
} else {
cfgFile = "Using config file: " + v.ConfigFileUsed()
func setupLogger(cfg *lib.Config) error {
loggerConfig := zap.NewProductionConfig()
loggerConfig.DisableCaller = true
if cfg.Debug {
loggerConfig.Level = zap.NewAtomicLevelAt(zap.DebugLevel)
}
loggerConfig.EncoderConfig.EncodeTime = zapcore.ISO8601TimeEncoder
loggerConfig.Encoding = cfg.LogFormat
logger, err := loggerConfig.Build()
if err != nil {
return err
}
zap.ReplaceGlobals(logger)
return nil
}