diff --git a/.env.example b/.env.example index 45d488e1..91f0b163 100644 --- a/.env.example +++ b/.env.example @@ -34,7 +34,7 @@ TINYAUTH_RESOURCES_PATH="./resources" TINYAUTH_SERVER_PORT=3000 # The address on which the server listens. TINYAUTH_SERVER_ADDRESS="0.0.0.0" -# The path to the Unix socket. +# Path of the Unix socket to listen on instead of TCP. An existing socket at this path is replaced, any other file is left untouched. TINYAUTH_SERVER_SOCKETPATH= # auth config diff --git a/internal/bootstrap/router_bootstrap.go b/internal/bootstrap/router_bootstrap.go index ae82c2d3..d31f4511 100644 --- a/internal/bootstrap/router_bootstrap.go +++ b/internal/bootstrap/router_bootstrap.go @@ -6,12 +6,12 @@ import ( "fmt" "net" "net/http" - "os" "time" "github.com/tinyauthapp/tinyauth/internal/controller" "github.com/tinyauthapp/tinyauth/internal/middleware" "github.com/tinyauthapp/tinyauth/internal/model" + "github.com/tinyauthapp/tinyauth/internal/utils" "go.uber.org/dig" "github.com/gin-gonic/gin" @@ -166,15 +166,16 @@ func (app *BootstrapApp) serveHTTP(ctx context.Context) error { } func (app *BootstrapApp) serveUnix(ctx context.Context) error { - _, err := os.Stat(app.config.Server.SocketPath) + removed, inUse, err := utils.RemoveExistingSocket(app.config.Server.SocketPath) - if err == nil { - app.log.App.Info().Msgf("Removing existing socket file %s", app.config.Server.SocketPath) - err := os.Remove(app.config.Server.SocketPath) + if err != nil { + return fmt.Errorf("failed to remove existing socket file: %w", err) + } - if err != nil { - return fmt.Errorf("failed to remove existing socket file: %w", err) - } + if inUse { + app.log.App.Warn().Msgf("Replaced socket %s that was still in use by another process, the socket path is where Tinyauth listens and should not be shared with other services", app.config.Server.SocketPath) + } else if removed { + app.log.App.Info().Msgf("Removed existing socket file %s", app.config.Server.SocketPath) } app.log.App.Info().Msgf("Starting server on unix socket %s", app.config.Server.SocketPath) diff --git a/internal/model/config.go b/internal/model/config.go index 9d514390..2ce25ecc 100644 --- a/internal/model/config.go +++ b/internal/model/config.go @@ -137,7 +137,7 @@ type ResourcesConfig struct { type ServerConfig struct { Port int `description:"The port on which the server listens." yaml:"port,omitempty"` Address string `description:"The address on which the server listens." yaml:"address,omitempty"` - SocketPath string `description:"The path to the Unix socket." yaml:"socketPath,omitempty"` + SocketPath string `description:"Path of the Unix socket to listen on instead of TCP. An existing socket at this path is replaced, any other file is left untouched." yaml:"socketPath,omitempty"` } type AuthConfig struct { diff --git a/internal/utils/fs_utils.go b/internal/utils/fs_utils.go index 8b9f28bf..703d1ee6 100644 --- a/internal/utils/fs_utils.go +++ b/internal/utils/fs_utils.go @@ -1,6 +1,11 @@ package utils -import "os" +import ( + "fmt" + "net" + "os" + "time" +) func ReadFile(file string) (string, error) { _, err := os.Stat(file) @@ -15,3 +20,33 @@ func ReadFile(file string) (string, error) { return string(data), nil } + +// RemoveExistingSocket removes the unix socket left at path by a previous (or still running) instance so the server +// can listen on it again. It reports whether a socket was removed and whether it was still accepting connections. +// Anything that is not a socket (e.g. a regular file or a directory) is never removed. +func RemoveExistingSocket(path string) (removed bool, inUse bool, err error) { + // stat follows symlinks, a symlink to a socket is replaced like the socket itself + info, err := os.Stat(path) + if os.IsNotExist(err) { + return false, false, nil + } + if err != nil { + return false, false, err + } + + if info.Mode().Type() != os.ModeSocket { + return false, false, fmt.Errorf("refusing to remove %s, it is not a unix socket", path) + } + + conn, err := net.DialTimeout("unix", path, time.Second) + if err == nil { + conn.Close() + inUse = true + } + + if err := os.Remove(path); err != nil { + return false, inUse, err + } + + return true, inUse, nil +} diff --git a/internal/utils/fs_utils_test.go b/internal/utils/fs_utils_test.go index 68154419..4326efb2 100644 --- a/internal/utils/fs_utils_test.go +++ b/internal/utils/fs_utils_test.go @@ -1,7 +1,9 @@ package utils import ( + "net" "os" + "path/filepath" "testing" "github.com/stretchr/testify/assert" @@ -30,3 +32,74 @@ func TestReadFile(t *testing.T) { assert.ErrorContains(t, err, "no such file or directory") assert.Equal(t, "", content) } + +func TestRemoveExistingSocket(t *testing.T) { + // Short directory, unix socket paths are limited to ~108 bytes + dir, err := os.MkdirTemp("", "ta") + require.NoError(t, err) + defer os.RemoveAll(dir) + + // Non-existing path + removed, inUse, err := RemoveExistingSocket(filepath.Join(dir, "missing.sock")) + assert.NoError(t, err) + assert.False(t, removed) + assert.False(t, inUse) + + // Regular file is never removed + file := filepath.Join(dir, "regular") + require.NoError(t, os.WriteFile(file, []byte("data"), 0600)) + removed, _, err = RemoveExistingSocket(file) + assert.ErrorContains(t, err, "not a unix socket") + assert.False(t, removed) + assert.FileExists(t, file) + + // Symlink to a regular file is never removed + fileLink := filepath.Join(dir, "regular.link") + require.NoError(t, os.Symlink(file, fileLink)) + removed, _, err = RemoveExistingSocket(fileLink) + assert.ErrorContains(t, err, "not a unix socket") + assert.False(t, removed) + assert.FileExists(t, fileLink) + + // Directory is never removed + removed, _, err = RemoveExistingSocket(dir) + assert.ErrorContains(t, err, "not a unix socket") + assert.False(t, removed) + + // Socket still in use (e.g. the previous instance during a rolling update) is replaced and reported + live := filepath.Join(dir, "live.sock") + listener, err := net.Listen("unix", live) + require.NoError(t, err) + listener.(*net.UnixListener).SetUnlinkOnClose(false) + defer listener.Close() + removed, inUse, err = RemoveExistingSocket(live) + assert.NoError(t, err) + assert.True(t, removed) + assert.True(t, inUse) + assert.NoFileExists(t, live) + + // Stale socket left behind by a previous run is removed + stale := filepath.Join(dir, "stale.sock") + staleListener, err := net.Listen("unix", stale) + require.NoError(t, err) + staleListener.(*net.UnixListener).SetUnlinkOnClose(false) + require.NoError(t, staleListener.Close()) + removed, inUse, err = RemoveExistingSocket(stale) + assert.NoError(t, err) + assert.True(t, removed) + assert.False(t, inUse) + assert.NoFileExists(t, stale) + + // Symlink to a socket is replaced like the socket itself + target := filepath.Join(dir, "target.sock") + targetListener, err := net.Listen("unix", target) + require.NoError(t, err) + targetListener.(*net.UnixListener).SetUnlinkOnClose(false) + require.NoError(t, targetListener.Close()) + link := filepath.Join(dir, "link.sock") + require.NoError(t, os.Symlink(target, link)) + removed, _, err = RemoveExistingSocket(link) + assert.NoError(t, err) + assert.True(t, removed) + assert.NoFileExists(t, link) +}