feat: serve DWH authentication over Unix socket
This commit is contained in:
@@ -0,0 +1,344 @@
|
||||
//go:build linux
|
||||
|
||||
// Package service provides the fail-closed Unix-socket HTTP verifier.
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/credential"
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/record"
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/registry"
|
||||
)
|
||||
|
||||
const verifiedKeyIDHeaderName = "X-DWH-Key-ID"
|
||||
|
||||
// Config binds the verifier to one protected registry and Unix socket.
|
||||
type Config struct {
|
||||
RegistryRoot string
|
||||
Socket string
|
||||
Logger *log.Logger
|
||||
Now func() time.Time
|
||||
}
|
||||
|
||||
var socketOwner = func(path string) (uint32, error) {
|
||||
info, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if info.Mode()&os.ModeSocket == 0 {
|
||||
return 0, fmt.Errorf("socket path is not a socket")
|
||||
}
|
||||
stat, ok := info.Sys().(*syscall.Stat_t)
|
||||
if !ok {
|
||||
return 0, fmt.Errorf("socket owner is unavailable")
|
||||
}
|
||||
return stat.Uid, nil
|
||||
}
|
||||
|
||||
// New returns the HTTP verifier. It has no network listener and is safe to use
|
||||
// with an httptest server only for synthetic test registries.
|
||||
func New(store *registry.Store, logger *log.Logger, now func() time.Time) http.Handler {
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
return http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
|
||||
if request.URL.Path != "/verify" {
|
||||
response.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
if request.Method != http.MethodGet {
|
||||
response.WriteHeader(http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
if store == nil || store.Check() != nil {
|
||||
writeDecision(logger, now, "unavailable", "")
|
||||
response.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
values := request.Header.Values("X-API-Key")
|
||||
if len(values) != 1 || len(values[0]) == 0 || len(values[0]) > credential.MaxHeaderBytes {
|
||||
writeDecision(logger, now, "deny_header", "")
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
value := []byte(values[0])
|
||||
if strings.HasPrefix(values[0], credential.Prefix+".") {
|
||||
keyID, ok := parseV1(value)
|
||||
if !ok {
|
||||
writeDecision(logger, now, "deny_v1", "")
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
verifyV1(response, store, logger, now, value, keyID)
|
||||
return
|
||||
}
|
||||
verifyLegacy(response, store, logger, now, value)
|
||||
})
|
||||
}
|
||||
|
||||
// ListenAndServe validates the complete registry before publishing a local
|
||||
// Unix listener. It never removes a non-socket collision.
|
||||
func ListenAndServe(ctx context.Context, config Config) error {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
if !canonicalAbsolute(config.RegistryRoot) || !canonicalAbsolute(config.Socket) {
|
||||
return errors.New("registry root and socket must be canonical absolute paths")
|
||||
}
|
||||
if err := validateSocketParent(config.Socket); err != nil {
|
||||
return err
|
||||
}
|
||||
store, err := registry.Open(config.RegistryRoot)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer store.Close()
|
||||
if err := store.Check(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := reclaimOwnedStaleSocket(config.Socket); err != nil {
|
||||
return err
|
||||
}
|
||||
listener, err := net.ListenUnix("unix", &net.UnixAddr{Name: config.Socket, Net: "unix"})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
listener.SetUnlinkOnClose(false)
|
||||
cleanup := func() error {
|
||||
if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||
return err
|
||||
}
|
||||
return removeOwnedSocket(config.Socket)
|
||||
}
|
||||
if err := os.Chmod(config.Socket, 0o660); err != nil {
|
||||
_ = cleanup()
|
||||
return err
|
||||
}
|
||||
server := &http.Server{Handler: New(store, config.Logger, config.Now)}
|
||||
serveResult := make(chan error, 1)
|
||||
go func() { serveResult <- server.Serve(listener) }()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
shutdownErr := server.Shutdown(shutdownCtx)
|
||||
cancel()
|
||||
serveErr := <-serveResult
|
||||
cleanupErr := removeOwnedSocket(config.Socket)
|
||||
if shutdownErr != nil {
|
||||
return shutdownErr
|
||||
}
|
||||
if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) {
|
||||
return serveErr
|
||||
}
|
||||
return cleanupErr
|
||||
case serveErr := <-serveResult:
|
||||
cleanupErr := removeOwnedSocket(config.Socket)
|
||||
if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) {
|
||||
return serveErr
|
||||
}
|
||||
return cleanupErr
|
||||
}
|
||||
}
|
||||
|
||||
func canonicalAbsolute(path string) bool { return filepath.IsAbs(path) && filepath.Clean(path) == path }
|
||||
|
||||
func validateSocketParent(socket string) error {
|
||||
parent := filepath.Dir(socket)
|
||||
info, err := os.Lstat(parent)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
|
||||
return errors.New("socket parent is not a directory")
|
||||
}
|
||||
resolved, err := filepath.EvalSymlinks(parent)
|
||||
if err != nil || resolved != parent {
|
||||
return errors.New("socket parent is not canonical")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func reclaimOwnedStaleSocket(socket string) error {
|
||||
state, err := ownedSocketState(socket)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
connection, err := net.DialTimeout("unix", socket, 100*time.Millisecond)
|
||||
if err == nil {
|
||||
_ = connection.Close()
|
||||
return errors.New("socket is already serving")
|
||||
}
|
||||
if !errors.Is(err, syscall.ECONNREFUSED) {
|
||||
return fmt.Errorf("cannot establish stale socket: %w", err)
|
||||
}
|
||||
// Revalidate type, owner, device, and inode immediately before unlinking.
|
||||
return removeSocketIfUnchanged(socket, state)
|
||||
}
|
||||
|
||||
func removeOwnedSocket(socket string) error {
|
||||
state, err := ownedSocketState(socket)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return removeSocketIfUnchanged(socket, state)
|
||||
}
|
||||
|
||||
type socketState struct {
|
||||
dev uint64
|
||||
ino uint64
|
||||
uid uint32
|
||||
}
|
||||
|
||||
func ownedSocketState(socket string) (socketState, error) {
|
||||
info, err := os.Lstat(socket)
|
||||
if err != nil {
|
||||
return socketState{}, err
|
||||
}
|
||||
if info.Mode()&os.ModeSocket == 0 {
|
||||
return socketState{}, errors.New("socket path collision")
|
||||
}
|
||||
stat, ok := info.Sys().(*syscall.Stat_t)
|
||||
if !ok {
|
||||
return socketState{}, errors.New("socket identity is unavailable")
|
||||
}
|
||||
owner, err := socketOwner(socket)
|
||||
if err != nil {
|
||||
return socketState{}, err
|
||||
}
|
||||
if owner != stat.Uid || owner != uint32(os.Geteuid()) {
|
||||
return socketState{}, errors.New("socket owner mismatch")
|
||||
}
|
||||
return socketState{dev: stat.Dev, ino: stat.Ino, uid: stat.Uid}, nil
|
||||
}
|
||||
|
||||
func removeSocketIfUnchanged(socket string, expected socketState) error {
|
||||
current, err := ownedSocketState(socket)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if current != expected {
|
||||
return errors.New("socket path changed")
|
||||
}
|
||||
return os.Remove(socket)
|
||||
}
|
||||
|
||||
func verifyV1(response http.ResponseWriter, store *registry.Store, logger *log.Logger, now func() time.Time, value []byte, keyID string) {
|
||||
if store == nil {
|
||||
writeDecision(logger, now, "unavailable", keyID)
|
||||
response.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
found, err := store.Find(keyID)
|
||||
if err != nil {
|
||||
if errors.Is(err, registry.ErrNotFound) {
|
||||
writeDecision(logger, now, "deny_unknown", keyID)
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
if errors.Is(err, registry.ErrRevoked) {
|
||||
writeDecision(logger, now, "deny_revoked", keyID)
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
writeDecision(logger, now, "unavailable", keyID)
|
||||
response.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
if _, ok := credential.VerifyV1(value, digest(found.SecretSHA256)); !ok {
|
||||
writeDecision(logger, now, "deny_mismatch", keyID)
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
writeDecision(logger, now, "allow_v1", keyID)
|
||||
response.Header().Set(verifiedKeyIDHeaderName, keyID)
|
||||
response.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func verifyLegacy(response http.ResponseWriter, store *registry.Store, logger *log.Logger, now func() time.Time, value []byte) {
|
||||
if store == nil {
|
||||
writeDecision(logger, now, "unavailable", "")
|
||||
response.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
found, err := store.FindLegacy()
|
||||
if err != nil {
|
||||
if errors.Is(err, registry.ErrNotFound) {
|
||||
writeDecision(logger, now, "deny_unknown", "")
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
if errors.Is(err, registry.ErrRevoked) {
|
||||
writeDecision(logger, now, "deny_revoked", "")
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
writeDecision(logger, now, "unavailable", "")
|
||||
response.WriteHeader(http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
if !credential.VerifyLegacy(value, digest(found.SecretSHA256)) {
|
||||
writeDecision(logger, now, "deny_mismatch", "")
|
||||
response.WriteHeader(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
writeDecision(logger, now, "allow_legacy", record.LegacyKeyID)
|
||||
response.Header().Set(verifiedKeyIDHeaderName, record.LegacyKeyID)
|
||||
response.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func parseV1(value []byte) (string, bool) {
|
||||
parts := strings.Split(string(value), ".")
|
||||
if len(parts) != 3 || parts[0] != credential.Prefix || len(parts[1]) != credential.KeyIDEncodedLength || len(parts[2]) != credential.SecretEncodedLength {
|
||||
return "", false
|
||||
}
|
||||
keyBytes, err := base64.RawURLEncoding.DecodeString(parts[1])
|
||||
if err != nil || len(keyBytes) != 12 || base64.RawURLEncoding.EncodeToString(keyBytes) != parts[1] {
|
||||
return "", false
|
||||
}
|
||||
secretBytes, err := base64.RawURLEncoding.DecodeString(parts[2])
|
||||
if err != nil || len(secretBytes) != 32 || base64.RawURLEncoding.EncodeToString(secretBytes) != parts[2] {
|
||||
return "", false
|
||||
}
|
||||
return parts[1], true
|
||||
}
|
||||
|
||||
func digest(encoded string) record.Digest {
|
||||
decoded, err := base64.RawURLEncoding.DecodeString(encoded)
|
||||
if err != nil || len(decoded) != sha256.Size {
|
||||
return record.Digest{}
|
||||
}
|
||||
var value record.Digest
|
||||
copy(value[:], decoded)
|
||||
return value
|
||||
}
|
||||
|
||||
func writeDecision(logger *log.Logger, now func() time.Time, decision, keyID string) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
timestamp := now().UTC().Format(time.RFC3339Nano)
|
||||
if keyID == "" {
|
||||
logger.Printf("timestamp=%s decision=%s", timestamp, decision)
|
||||
return
|
||||
}
|
||||
logger.Printf("timestamp=%s decision=%s key_id=%s", timestamp, decision, keyID)
|
||||
}
|
||||
Reference in New Issue
Block a user