371 lines
10 KiB
Go
371 lines
10 KiB
Go
//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
|
|
}
|
|
|
|
var socketParentOwner = func(path string) (uint32, error) {
|
|
info, err := os.Lstat(path)
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
stat, ok := info.Sys().(*syscall.Stat_t)
|
|
if !ok {
|
|
return 0, fmt.Errorf("socket parent 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.OpenReadOnly(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")
|
|
}
|
|
if info.Mode().Perm()&0o022 != 0 {
|
|
return errors.New("socket parent is group or world writable")
|
|
}
|
|
stat, ok := info.Sys().(*syscall.Stat_t)
|
|
if !ok {
|
|
return errors.New("socket parent owner is unavailable")
|
|
}
|
|
owner, err := socketParentOwner(parent)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if stat.Uid != uint32(os.Geteuid()) || owner != stat.Uid {
|
|
return errors.New("socket parent owner mismatch")
|
|
}
|
|
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)
|
|
}
|