Files

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)
}