mirror of
https://github.com/hacdias/webdav.git
synced 2026-09-22 03:20:41 +08:00
refactor: code cleanup, stricter config validation (#155)
This commit is contained in:
+153
@@ -0,0 +1,153 @@
|
||||
package lib
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/spf13/pflag"
|
||||
"github.com/spf13/viper"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Permissions `mapstructure:",squash"`
|
||||
Debug bool
|
||||
Address string
|
||||
Port int
|
||||
TLS bool
|
||||
Cert string
|
||||
Key string
|
||||
Prefix string
|
||||
NoSniff bool
|
||||
LogFormat string
|
||||
Auth bool
|
||||
CORS CORS
|
||||
Users []User
|
||||
}
|
||||
|
||||
func ParseConfig(filename string, flags *pflag.FlagSet) (*Config, error) {
|
||||
v := viper.New()
|
||||
|
||||
// Configure flags bindings
|
||||
if flags != nil {
|
||||
err := v.BindPFlags(flags)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
err = v.BindPFlag("LogFormat", flags.Lookup("log_format"))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// Configuration file settings
|
||||
v.AddConfigPath(".")
|
||||
v.AddConfigPath("/etc/webdav/")
|
||||
v.SetConfigName("config")
|
||||
if filename != "" {
|
||||
v.SetConfigFile(filename)
|
||||
}
|
||||
|
||||
// Environment settings
|
||||
v.SetEnvPrefix("wd")
|
||||
v.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
|
||||
v.AutomaticEnv()
|
||||
|
||||
// Defaults
|
||||
v.SetDefault("CORS.AllowedHeaders", []string{"*"})
|
||||
v.SetDefault("CORS.AllowedHosts", []string{"*"})
|
||||
v.SetDefault("CORS.AllowedMethods", []string{"*"})
|
||||
|
||||
// Read and unmarshal configuration
|
||||
err := v.ReadInConfig()
|
||||
if err != nil {
|
||||
if _, ok := err.(viper.ConfigFileNotFoundError); !ok {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
cfg := &Config{}
|
||||
err = v.Unmarshal(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Cascade user settings
|
||||
for i := range cfg.Users {
|
||||
if !v.IsSet(fmt.Sprintf("Users.%d.Scope", i)) {
|
||||
cfg.Users[i].Scope = cfg.Scope
|
||||
}
|
||||
|
||||
if !v.IsSet(fmt.Sprintf("Users.%d.Modify", i)) {
|
||||
cfg.Users[i].Modify = cfg.Modify
|
||||
}
|
||||
|
||||
if !v.IsSet(fmt.Sprintf("Users.%d.Rules", i)) {
|
||||
cfg.Users[i].Rules = cfg.Rules
|
||||
}
|
||||
}
|
||||
|
||||
err = cfg.Validate()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func (c *Config) Validate() error {
|
||||
var err error
|
||||
|
||||
if c.Auth && len(c.Users) == 0 {
|
||||
return errors.New("invalid config: auth cannot be enabled without users")
|
||||
}
|
||||
|
||||
if !c.Auth && len(c.Users) != 0 {
|
||||
return errors.New("invalid config: auth cannot be disabled with users defined")
|
||||
}
|
||||
|
||||
if c.TLS {
|
||||
if c.Cert == "" {
|
||||
return errors.New("invalid config: Cert must be defined if TLS is activated")
|
||||
}
|
||||
|
||||
if c.Key == "" {
|
||||
return errors.New("invalid config: Key must be defined if TLS is activated")
|
||||
}
|
||||
|
||||
c.Cert, err = filepath.Abs(c.Cert)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid config: %w", err)
|
||||
}
|
||||
|
||||
c.Key, err = filepath.Abs(c.Key)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid config: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
err = c.Permissions.Validate()
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid config: %w", err)
|
||||
}
|
||||
|
||||
for _, u := range c.Users {
|
||||
err := u.Validate()
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid config: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type CORS struct {
|
||||
Enabled bool
|
||||
Credentials bool
|
||||
AllowedHeaders []string
|
||||
AllowedHosts []string
|
||||
AllowedMethods []string
|
||||
ExposedHeaders []string
|
||||
}
|
||||
@@ -0,0 +1,71 @@
|
||||
package lib
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func writeAndParseConfig(t *testing.T, content string) *Config {
|
||||
tmpDir := t.TempDir()
|
||||
tmpFile := filepath.Join(tmpDir, "config.yml")
|
||||
|
||||
err := os.WriteFile(tmpFile, []byte(content), 0666)
|
||||
require.NoError(t, err)
|
||||
|
||||
cfg, err := ParseConfig(tmpFile, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
return cfg
|
||||
}
|
||||
|
||||
func TestConfigDefaults(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cfg := writeAndParseConfig(t, "")
|
||||
require.NoError(t, cfg.Validate())
|
||||
|
||||
require.EqualValues(t, []string{"*"}, cfg.CORS.AllowedHeaders)
|
||||
require.EqualValues(t, []string{"*"}, cfg.CORS.AllowedHosts)
|
||||
require.EqualValues(t, []string{"*"}, cfg.CORS.AllowedMethods)
|
||||
}
|
||||
|
||||
func TestConfigCascade(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
content := `
|
||||
auth: true
|
||||
scope: /
|
||||
modify: true
|
||||
rules:
|
||||
- path: /public/access/
|
||||
modify: true
|
||||
|
||||
users:
|
||||
- username: admin
|
||||
password: admin
|
||||
- username: basic
|
||||
password: basic
|
||||
scope: /basic
|
||||
modify: false
|
||||
rules: []`
|
||||
|
||||
cfg := writeAndParseConfig(t, content)
|
||||
require.NoError(t, cfg.Validate())
|
||||
|
||||
require.True(t, cfg.Modify)
|
||||
require.Equal(t, "/", cfg.Scope)
|
||||
require.Len(t, cfg.Rules, 1)
|
||||
|
||||
require.Len(t, cfg.Users, 2)
|
||||
|
||||
require.True(t, cfg.Users[0].Modify)
|
||||
require.Equal(t, "/", cfg.Users[0].Scope)
|
||||
require.Len(t, cfg.Users[0].Rules, 1)
|
||||
|
||||
require.False(t, cfg.Users[1].Modify)
|
||||
require.Equal(t, "/basic", cfg.Users[1].Scope)
|
||||
require.Len(t, cfg.Users[1].Rules, 0)
|
||||
}
|
||||
+40
-41
@@ -9,12 +9,44 @@ import (
|
||||
"golang.org/x/net/webdav"
|
||||
)
|
||||
|
||||
// NoSniffFileInfo wraps any generic FileInfo interface and bypasses mime type sniffing.
|
||||
type NoSniffFileInfo struct {
|
||||
type Dir struct {
|
||||
webdav.Dir
|
||||
noSniff bool
|
||||
}
|
||||
|
||||
func (d Dir) Stat(ctx context.Context, name string) (os.FileInfo, error) {
|
||||
// Skip wrapping if NoSniff is off
|
||||
if !d.noSniff {
|
||||
return d.Dir.Stat(ctx, name)
|
||||
}
|
||||
|
||||
info, err := d.Dir.Stat(ctx, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return noSniffFileInfo{info}, nil
|
||||
}
|
||||
|
||||
func (d Dir) OpenFile(ctx context.Context, name string, flag int, perm os.FileMode) (webdav.File, error) {
|
||||
// Skip wrapping if NoSniff is off
|
||||
if !d.noSniff {
|
||||
return d.Dir.OpenFile(ctx, name, flag, perm)
|
||||
}
|
||||
|
||||
file, err := d.Dir.OpenFile(ctx, name, flag, perm)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return noSniffFile{File: file}, nil
|
||||
}
|
||||
|
||||
type noSniffFileInfo struct {
|
||||
os.FileInfo
|
||||
}
|
||||
|
||||
func (w NoSniffFileInfo) ContentType(ctx context.Context) (contentType string, err error) {
|
||||
func (w noSniffFileInfo) ContentType(ctx context.Context) (contentType string, err error) {
|
||||
if mimeType := mime.TypeByExtension(path.Ext(w.FileInfo.Name())); mimeType != "" {
|
||||
// We can figure out the mime from the extension.
|
||||
return mimeType, nil
|
||||
@@ -24,60 +56,27 @@ func (w NoSniffFileInfo) ContentType(ctx context.Context) (contentType string, e
|
||||
}
|
||||
}
|
||||
|
||||
type WebDavDir struct {
|
||||
webdav.Dir
|
||||
NoSniff bool
|
||||
}
|
||||
|
||||
func (d WebDavDir) Stat(ctx context.Context, name string) (os.FileInfo, error) {
|
||||
// Skip wrapping if NoSniff is off
|
||||
if !d.NoSniff {
|
||||
return d.Dir.Stat(ctx, name)
|
||||
}
|
||||
|
||||
info, err := d.Dir.Stat(ctx, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return NoSniffFileInfo{info}, nil
|
||||
}
|
||||
|
||||
func (d WebDavDir) OpenFile(ctx context.Context, name string, flag int, perm os.FileMode) (webdav.File, error) {
|
||||
// Skip wrapping if NoSniff is off
|
||||
if !d.NoSniff {
|
||||
return d.Dir.OpenFile(ctx, name, flag, perm)
|
||||
}
|
||||
|
||||
file, err := d.Dir.OpenFile(ctx, name, flag, perm)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return WebDavFile{File: file}, nil
|
||||
}
|
||||
|
||||
type WebDavFile struct {
|
||||
type noSniffFile struct {
|
||||
webdav.File
|
||||
}
|
||||
|
||||
func (f WebDavFile) Stat() (os.FileInfo, error) {
|
||||
func (f noSniffFile) Stat() (os.FileInfo, error) {
|
||||
info, err := f.File.Stat()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return NoSniffFileInfo{info}, nil
|
||||
return noSniffFileInfo{info}, nil
|
||||
}
|
||||
|
||||
func (f WebDavFile) Readdir(count int) (fis []os.FileInfo, err error) {
|
||||
func (f noSniffFile) Readdir(count int) (fis []os.FileInfo, err error) {
|
||||
fis, err = f.File.Readdir(count)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for i := range fis {
|
||||
fis[i] = NoSniffFileInfo{fis[i]}
|
||||
fis[i] = noSniffFileInfo{fis[i]}
|
||||
}
|
||||
return fis, nil
|
||||
}
|
||||
+133
@@ -0,0 +1,133 @@
|
||||
package lib
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/rs/cors"
|
||||
"go.uber.org/zap"
|
||||
"golang.org/x/net/webdav"
|
||||
)
|
||||
|
||||
type handlerUser struct {
|
||||
User
|
||||
webdav.Handler
|
||||
}
|
||||
|
||||
type Handler struct {
|
||||
*Config
|
||||
user *handlerUser
|
||||
users map[string]*handlerUser
|
||||
}
|
||||
|
||||
func NewHandler(c *Config) (http.Handler, error) {
|
||||
h := &Handler{
|
||||
user: &handlerUser{
|
||||
User: User{
|
||||
Permissions: c.Permissions,
|
||||
},
|
||||
Handler: webdav.Handler{
|
||||
Prefix: c.Prefix,
|
||||
FileSystem: Dir{
|
||||
Dir: webdav.Dir(c.Scope),
|
||||
noSniff: c.NoSniff,
|
||||
},
|
||||
LockSystem: webdav.NewMemLS(),
|
||||
},
|
||||
},
|
||||
users: map[string]*handlerUser{},
|
||||
}
|
||||
|
||||
for _, u := range c.Users {
|
||||
h.users[u.Username] = &handlerUser{
|
||||
User: u,
|
||||
Handler: webdav.Handler{
|
||||
Prefix: c.Prefix,
|
||||
FileSystem: Dir{
|
||||
Dir: webdav.Dir(u.Scope),
|
||||
noSniff: c.NoSniff,
|
||||
},
|
||||
LockSystem: webdav.NewMemLS(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
if c.CORS.Enabled {
|
||||
return cors.New(cors.Options{
|
||||
AllowCredentials: c.CORS.Credentials,
|
||||
AllowedOrigins: c.CORS.AllowedHosts,
|
||||
AllowedMethods: c.CORS.AllowedMethods,
|
||||
AllowedHeaders: c.CORS.AllowedHeaders,
|
||||
OptionsPassthrough: false,
|
||||
}).Handler(h), nil
|
||||
}
|
||||
|
||||
return h, nil
|
||||
}
|
||||
|
||||
// ServeHTTP determines if the request is for this plugin, and if all prerequisites are met.
|
||||
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
user := h.user
|
||||
|
||||
// Authentication
|
||||
if h.Auth {
|
||||
w.Header().Set("WWW-Authenticate", `Basic realm="Restricted"`)
|
||||
|
||||
// Gets the correct user for this request.
|
||||
username, password, ok := r.BasicAuth()
|
||||
zap.L().Info("login attempt", zap.String("username", username), zap.String("remote_address", r.RemoteAddr))
|
||||
if !ok {
|
||||
http.Error(w, "Not authorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
user, ok = h.users[username]
|
||||
if !ok {
|
||||
http.Error(w, "Not authorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
if !user.checkPassword(password) {
|
||||
zap.L().Info("invalid password", zap.String("username", username), zap.String("remote_address", r.RemoteAddr))
|
||||
http.Error(w, "Not authorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
zap.L().Info("user authorized", zap.String("username", username))
|
||||
}
|
||||
|
||||
// Checks for user permissions relatively to this PATH.
|
||||
allowed := user.Allowed(r)
|
||||
|
||||
zap.L().Debug("allowed & method & path", zap.Bool("allowed", allowed), zap.String("method", r.Method), zap.String("path", r.URL.Path))
|
||||
|
||||
if !allowed {
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
if r.Method == "HEAD" {
|
||||
w = newResponseWriterNoBody(w)
|
||||
}
|
||||
|
||||
// Excerpt from RFC4918, section 9.4:
|
||||
//
|
||||
// GET, when applied to a collection, may return the contents of an
|
||||
// "index.html" resource, a human-readable view of the contents of
|
||||
// the collection, or something else altogether.
|
||||
//
|
||||
// Get, when applied to collection, will return the same as PROPFIND method.
|
||||
if r.Method == "GET" && strings.HasPrefix(r.URL.Path, user.Prefix) {
|
||||
info, err := user.FileSystem.Stat(r.Context(), strings.TrimPrefix(r.URL.Path, user.Prefix))
|
||||
if err == nil && info.IsDir() {
|
||||
r.Method = "PROPFIND"
|
||||
|
||||
if r.Header.Get("Depth") == "" {
|
||||
r.Header.Add("Depth", "1")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Runs the WebDAV.
|
||||
user.ServeHTTP(w, r)
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package lib
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var readMethods = []string{
|
||||
http.MethodGet,
|
||||
http.MethodHead,
|
||||
http.MethodOptions,
|
||||
"PROPFIND",
|
||||
}
|
||||
|
||||
type Rule struct {
|
||||
Regex bool
|
||||
Allow bool
|
||||
Modify bool
|
||||
Path string
|
||||
// TODO: remove Regex and replace by this. It encodes
|
||||
Regexp *regexp.Regexp `mapstructure:"-"`
|
||||
}
|
||||
|
||||
func (r *Rule) Validate() error {
|
||||
if r.Regex {
|
||||
rp, err := regexp.Compile(r.Path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid rule: %w", err)
|
||||
}
|
||||
r.Regexp = rp
|
||||
r.Path = ""
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Matches checks if [Rule] matches the given path.
|
||||
func (r *Rule) Matches(path string) bool {
|
||||
if r.Regex {
|
||||
return r.Regexp.MatchString(path)
|
||||
}
|
||||
|
||||
return strings.HasPrefix(path, r.Path)
|
||||
}
|
||||
|
||||
type Permissions struct {
|
||||
Scope string
|
||||
Modify bool
|
||||
Rules []*Rule
|
||||
}
|
||||
|
||||
// Allowed checks if the user has permission to access a directory/file
|
||||
func (p Permissions) Allowed(r *http.Request) bool {
|
||||
// Determine whether or not it is a read or write request.
|
||||
readRequest := false
|
||||
for _, method := range readMethods {
|
||||
if r.Method == method {
|
||||
readRequest = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Go through rules beginning from the last one.
|
||||
for i := len(p.Rules) - 1; i >= 0; i-- {
|
||||
rule := p.Rules[i]
|
||||
|
||||
if rule.Matches(r.URL.Path) {
|
||||
return rule.Allow && (readRequest || rule.Modify)
|
||||
}
|
||||
}
|
||||
|
||||
return readRequest || p.Modify
|
||||
}
|
||||
|
||||
func (p *Permissions) Validate() error {
|
||||
for _, r := range p.Rules {
|
||||
if err := r.Validate(); err != nil {
|
||||
return fmt.Errorf("invalid permissions: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package lib
|
||||
|
||||
import "net/http"
|
||||
|
||||
var _ http.ResponseWriter = responseWriterNoBody{}
|
||||
|
||||
// responseWriterNoBody is a wrapper used to suppress the body of the response
|
||||
// to a request. Mainly used for HEAD requests.
|
||||
type responseWriterNoBody struct {
|
||||
http.ResponseWriter
|
||||
}
|
||||
|
||||
// newResponseWriterNoBody creates a new responseWriterNoBody.
|
||||
func newResponseWriterNoBody(w http.ResponseWriter) *responseWriterNoBody {
|
||||
return &responseWriterNoBody{w}
|
||||
}
|
||||
|
||||
// Write suppress the body.
|
||||
func (w responseWriterNoBody) Write(data []byte) (int, error) {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// WriteHeader writes the header to the http.ResponseWriter.
|
||||
func (w responseWriterNoBody) WriteHeader(statusCode int) {
|
||||
w.Header().Del("Content-Length")
|
||||
w.ResponseWriter.WriteHeader(statusCode)
|
||||
}
|
||||
Executable → Regular
+39
-37
@@ -1,50 +1,52 @@
|
||||
package lib
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/net/webdav"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// Rule is a disallow/allow rule.
|
||||
type Rule struct {
|
||||
Regex bool
|
||||
Allow bool
|
||||
Modify bool
|
||||
Path string
|
||||
Regexp *regexp.Regexp
|
||||
}
|
||||
|
||||
// User contains the settings of each user.
|
||||
type User struct {
|
||||
Username string
|
||||
Password string
|
||||
Scope string
|
||||
Modify bool
|
||||
Rules []*Rule
|
||||
Handler *webdav.Handler
|
||||
Permissions `mapstructure:",squash"`
|
||||
Username string
|
||||
Password string
|
||||
}
|
||||
|
||||
// Allowed checks if the user has permission to access a directory/file
|
||||
func (u User) Allowed(url string, noModification bool) bool {
|
||||
var rule *Rule
|
||||
i := len(u.Rules) - 1
|
||||
|
||||
for i >= 0 {
|
||||
rule = u.Rules[i]
|
||||
|
||||
isAllowed := rule.Allow && (noModification || rule.Modify)
|
||||
if rule.Regex {
|
||||
if rule.Regexp.MatchString(url) {
|
||||
return isAllowed
|
||||
}
|
||||
} else if strings.HasPrefix(url, rule.Path) {
|
||||
return isAllowed
|
||||
}
|
||||
|
||||
i--
|
||||
func (u User) checkPassword(input string) bool {
|
||||
if strings.HasPrefix(u.Password, "{bcrypt}") {
|
||||
savedPassword := strings.TrimPrefix(u.Password, "{bcrypt}")
|
||||
return bcrypt.CompareHashAndPassword([]byte(savedPassword), []byte(input)) == nil
|
||||
}
|
||||
|
||||
return noModification || u.Modify
|
||||
return u.Password == input
|
||||
}
|
||||
|
||||
func (u *User) Validate() error {
|
||||
if u.Username == "" {
|
||||
return errors.New("invalid user: username must be set")
|
||||
}
|
||||
|
||||
if u.Password == "" {
|
||||
return fmt.Errorf("invalid user %q: password must be set", u.Username)
|
||||
} else if strings.HasPrefix(u.Password, "{env}") {
|
||||
|
||||
env := strings.TrimPrefix(u.Password, "{env}")
|
||||
if env == "" {
|
||||
return fmt.Errorf("invalid user %q: password environment variable not set", u.Username)
|
||||
}
|
||||
|
||||
u.Password = os.Getenv(env)
|
||||
if u.Password == "" {
|
||||
return fmt.Errorf("invalid user %q: password environment variable is empty", u.Username)
|
||||
}
|
||||
}
|
||||
|
||||
if err := u.Permissions.Validate(); err != nil {
|
||||
return fmt.Errorf("invalid user %q: %w", u.Username, err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,25 +0,0 @@
|
||||
package lib
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func checkPassword(saved, input string) bool {
|
||||
if strings.HasPrefix(saved, "{bcrypt}") {
|
||||
savedPassword := strings.TrimPrefix(saved, "{bcrypt}")
|
||||
return bcrypt.CompareHashAndPassword([]byte(savedPassword), []byte(input)) == nil
|
||||
}
|
||||
|
||||
return saved == input
|
||||
}
|
||||
|
||||
func isAllowedHost(allowedHosts []string, origin string) bool {
|
||||
for _, host := range allowedHosts {
|
||||
if host == origin {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
-174
@@ -1,174 +0,0 @@
|
||||
package lib
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// CorsCfg is the CORS config.
|
||||
type CorsCfg struct {
|
||||
Enabled bool
|
||||
Credentials bool
|
||||
AllowedHeaders []string
|
||||
AllowedHosts []string
|
||||
AllowedMethods []string
|
||||
ExposedHeaders []string
|
||||
}
|
||||
|
||||
// Config is the configuration of a WebDAV instance.
|
||||
type Config struct {
|
||||
*User
|
||||
Auth bool
|
||||
Debug bool
|
||||
NoSniff bool
|
||||
Cors CorsCfg
|
||||
Users map[string]*User
|
||||
LogFormat string
|
||||
}
|
||||
|
||||
// ServeHTTP determines if the request is for this plugin, and if all prerequisites are met.
|
||||
func (c *Config) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
u := c.User
|
||||
requestOrigin := r.Header.Get("Origin")
|
||||
|
||||
// Add CORS headers before any operation so even on a 401 unauthorized status, CORS will work.
|
||||
if c.Cors.Enabled && requestOrigin != "" {
|
||||
headers := w.Header()
|
||||
|
||||
allowedHeaders := strings.Join(c.Cors.AllowedHeaders, ", ")
|
||||
allowedMethods := strings.Join(c.Cors.AllowedMethods, ", ")
|
||||
exposedHeaders := strings.Join(c.Cors.ExposedHeaders, ", ")
|
||||
|
||||
allowAllHosts := len(c.Cors.AllowedHosts) == 1 && c.Cors.AllowedHosts[0] == "*"
|
||||
allowedHost := isAllowedHost(c.Cors.AllowedHosts, requestOrigin)
|
||||
|
||||
if allowAllHosts {
|
||||
headers.Set("Access-Control-Allow-Origin", "*")
|
||||
} else if allowedHost {
|
||||
headers.Set("Access-Control-Allow-Origin", requestOrigin)
|
||||
}
|
||||
|
||||
if allowAllHosts || allowedHost {
|
||||
headers.Set("Access-Control-Allow-Headers", allowedHeaders)
|
||||
headers.Set("Access-Control-Allow-Methods", allowedMethods)
|
||||
|
||||
if c.Cors.Credentials {
|
||||
headers.Set("Access-Control-Allow-Credentials", "true")
|
||||
}
|
||||
|
||||
if len(c.Cors.ExposedHeaders) > 0 {
|
||||
headers.Set("Access-Control-Expose-Headers", exposedHeaders)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if r.Method == "OPTIONS" && c.Cors.Enabled && requestOrigin != "" {
|
||||
return
|
||||
}
|
||||
|
||||
// Authentication
|
||||
if c.Auth {
|
||||
w.Header().Set("WWW-Authenticate", `Basic realm="Restricted"`)
|
||||
|
||||
// Gets the correct user for this request.
|
||||
username, password, ok := r.BasicAuth()
|
||||
zap.L().Info("login attempt", zap.String("username", username), zap.String("remote_address", r.RemoteAddr))
|
||||
if !ok {
|
||||
http.Error(w, "Not authorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
user, ok := c.Users[username]
|
||||
if !ok {
|
||||
http.Error(w, "Not authorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
if !checkPassword(user.Password, password) {
|
||||
zap.L().Info("invalid password", zap.String("username", username), zap.String("remote_address", r.RemoteAddr))
|
||||
http.Error(w, "Not authorized", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
u = user
|
||||
zap.L().Info("user authorized", zap.String("username", username))
|
||||
} else {
|
||||
// Even if Auth is disabled, we might want to get
|
||||
// the user from the Basic Auth header. Useful for Caddy
|
||||
// plugin implementation.
|
||||
username, _, ok := r.BasicAuth()
|
||||
if ok {
|
||||
if user, ok := c.Users[username]; ok {
|
||||
u = user
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Checks for user permissions relatively to this PATH.
|
||||
noModification := r.Method == "GET" || r.Method == "HEAD" ||
|
||||
r.Method == "OPTIONS" || r.Method == "PROPFIND"
|
||||
|
||||
allowed := u.Allowed(r.URL.Path, noModification)
|
||||
|
||||
zap.L().Debug("allowed & method & path", zap.Bool("allowed", allowed), zap.String("method", r.Method), zap.String("path", r.URL.Path))
|
||||
|
||||
if !allowed {
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
|
||||
if r.Method == "HEAD" {
|
||||
w = newResponseWriterNoBody(w)
|
||||
}
|
||||
|
||||
// Excerpt from RFC4918, section 9.4:
|
||||
//
|
||||
// GET, when applied to a collection, may return the contents of an
|
||||
// "index.html" resource, a human-readable view of the contents of
|
||||
// the collection, or something else altogether.
|
||||
//
|
||||
// Get, when applied to collection, will return the same as PROPFIND method.
|
||||
if r.Method == "GET" && strings.HasPrefix(r.URL.Path, u.Handler.Prefix) {
|
||||
info, err := u.Handler.FileSystem.Stat(context.TODO(), strings.TrimPrefix(r.URL.Path, u.Handler.Prefix))
|
||||
if err == nil && info.IsDir() {
|
||||
r.Method = "PROPFIND"
|
||||
|
||||
if r.Header.Get("Depth") == "" {
|
||||
r.Header.Add("Depth", "1")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Runs the WebDAV.
|
||||
//u.Handler.LockSystem = webdav.NewMemLS()
|
||||
u.Handler.ServeHTTP(w, r)
|
||||
}
|
||||
|
||||
// responseWriterNoBody is a wrapper used to suprress the body of the response
|
||||
// to a request. Mainly used for HEAD requests.
|
||||
type responseWriterNoBody struct {
|
||||
http.ResponseWriter
|
||||
}
|
||||
|
||||
// newResponseWriterNoBody creates a new responseWriterNoBody.
|
||||
func newResponseWriterNoBody(w http.ResponseWriter) *responseWriterNoBody {
|
||||
return &responseWriterNoBody{w}
|
||||
}
|
||||
|
||||
// Header executes the Header method from the http.ResponseWriter.
|
||||
func (w responseWriterNoBody) Header() http.Header {
|
||||
return w.ResponseWriter.Header()
|
||||
}
|
||||
|
||||
// Write suprresses the body.
|
||||
func (w responseWriterNoBody) Write(data []byte) (int, error) {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
// WriteHeader writes the header to the http.ResponseWriter.
|
||||
func (w responseWriterNoBody) WriteHeader(statusCode int) {
|
||||
w.ResponseWriter.WriteHeader(statusCode)
|
||||
}
|
||||
Reference in New Issue
Block a user