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