feat: serve DWH authentication over Unix socket
This commit is contained in:
@@ -14,6 +14,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -25,15 +26,15 @@ import (
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/record"
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/registry"
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/securefile"
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/service"
|
||||
)
|
||||
|
||||
const (
|
||||
exitOK = 0
|
||||
exitUsage = 2
|
||||
exitNotFound = 3
|
||||
exitIntegrity = 4
|
||||
maxReasonBytes = 160
|
||||
commandServeMsg = "serve is not available in this release"
|
||||
exitOK = 0
|
||||
exitUsage = 2
|
||||
exitNotFound = 3
|
||||
exitIntegrity = 4
|
||||
maxReasonBytes = 160
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -41,6 +42,7 @@ var (
|
||||
generateCredential = func() (credential.Material, error) { return credential.Generate(rand.Reader) }
|
||||
addRecord = func(store *registry.Store, value record.Record) error { return store.Add(value) }
|
||||
closeStore = func(store *registry.Store) error { return store.Close() }
|
||||
listenAndServe = service.ListenAndServe
|
||||
)
|
||||
|
||||
// Run executes one dwh-auth invocation and returns its stable process exit
|
||||
@@ -53,7 +55,7 @@ func Run(ctx context.Context, args []string, _ io.Reader, stdout, stderr io.Writ
|
||||
return exitIntegrity
|
||||
}
|
||||
if len(args) > 0 && args[0] == "serve" {
|
||||
return runServe(args[1:], stderr)
|
||||
return runServe(ctx, args[1:], stderr)
|
||||
}
|
||||
root, rest, err := parseRoot(args)
|
||||
if err != nil {
|
||||
@@ -437,17 +439,21 @@ func runCheck(ctx context.Context, root string, args []string, stdout, stderr io
|
||||
return exitOK
|
||||
}
|
||||
|
||||
func runServe(args []string, stderr io.Writer) int {
|
||||
if len(args) != 4 || args[0] != "--registry-root" || args[2] != "--socket" || !filepath.IsAbs(args[1]) || !filepath.IsAbs(args[3]) {
|
||||
func runServe(ctx context.Context, args []string, stderr io.Writer) int {
|
||||
if len(args) != 4 || args[0] != "--registry-root" || args[2] != "--socket" || !canonicalAbsolute(args[1]) || !canonicalAbsolute(args[3]) {
|
||||
writeLine(stderr, "unsafe invocation")
|
||||
return exitUsage
|
||||
}
|
||||
// The Unix-socket service is intentionally introduced by Task 4. Keeping
|
||||
// this exact grammar reserved prevents accidental plaintext or network fallback.
|
||||
writeLine(stderr, commandServeMsg)
|
||||
return exitIntegrity
|
||||
logger := log.New(stderr, "", 0)
|
||||
if err := listenAndServe(ctx, service.Config{RegistryRoot: args[1], Socket: args[3], Logger: logger}); err != nil {
|
||||
writeLine(stderr, "integrity failure")
|
||||
return exitIntegrity
|
||||
}
|
||||
return exitOK
|
||||
}
|
||||
|
||||
func canonicalAbsolute(path string) bool { return filepath.IsAbs(path) && filepath.Clean(path) == path }
|
||||
|
||||
func snapshotContainsKey(store *registry.Store, keyID string) (bool, error) {
|
||||
values, err := store.List()
|
||||
if err != nil {
|
||||
|
||||
@@ -19,6 +19,7 @@ import (
|
||||
"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"
|
||||
"github.com/aritmolab/thothii/tools/dwh-auth/internal/service"
|
||||
)
|
||||
|
||||
const sentinelSecret = "SENTINEL-DWH-SECRET-must-never-be-emitted"
|
||||
@@ -295,17 +296,29 @@ func TestMetadataAndExpiryAreValidatedBeforeKeyGeneration(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeAcceptsOnlyExactTaskFourReservation(t *testing.T) {
|
||||
func TestServeAcceptsOnlyExactCheckedServiceConfig(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
socket := filepath.Join(t.TempDir(), "verify.sock")
|
||||
oldListen := listenAndServe
|
||||
t.Cleanup(func() { listenAndServe = oldListen })
|
||||
called := false
|
||||
listenAndServe = func(ctx context.Context, config service.Config) error {
|
||||
called = true
|
||||
if ctx == nil || config.RegistryRoot != root || config.Socket != socket || config.Logger == nil {
|
||||
t.Fatalf("serve config = %#v", config)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
stdout, stderr, code := run(t, "serve", "--registry-root", root, "--socket", socket)
|
||||
if code != 4 || stdout != "" || stderr != "serve is not available in this release\n" {
|
||||
t.Fatalf("valid serve reservation = (%d, %q, %q)", code, stdout, stderr)
|
||||
if code != 0 || stdout != "" || stderr != "" || !called {
|
||||
t.Fatalf("valid serve = (%d, %q, %q), called=%v", code, stdout, stderr, called)
|
||||
}
|
||||
invalid := [][]string{
|
||||
{"serve", "--registry-root", "relative", "--socket", socket},
|
||||
{"serve", "--registry-root", root + "/.", "--socket", socket},
|
||||
{"serve", "--registry-root", root},
|
||||
{"serve", "--registry-root", root, "--socket", "relative"},
|
||||
{"serve", "--registry-root", root, "--socket", filepath.Dir(socket) + "/./verify.sock"},
|
||||
{"serve", "--registry-root", root, "--socket", socket, "--socket", socket},
|
||||
{"serve", "--socket", socket, "--registry-root", root},
|
||||
{"serve", "--registry-root", root, "--socket", socket, "--unknown", "x"},
|
||||
@@ -319,6 +332,18 @@ func TestServeAcceptsOnlyExactTaskFourReservation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeSanitizesServiceStartupFailure(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
socket := filepath.Join(t.TempDir(), "verify.sock")
|
||||
oldListen := listenAndServe
|
||||
t.Cleanup(func() { listenAndServe = oldListen })
|
||||
listenAndServe = func(context.Context, service.Config) error { return errors.New(sentinelSecret) }
|
||||
stdout, stderr, code := run(t, "serve", "--registry-root", root, "--socket", socket)
|
||||
if code != 4 || stdout != "" || stderr != "integrity failure\n" || strings.Contains(stdout+stderr, sentinelSecret) {
|
||||
t.Fatalf("service startup failure = (%d, %q, %q)", code, stdout, stderr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateCleansOnlyWhenFailedPublicationProvesKeyAbsent(t *testing.T) {
|
||||
oldAdd, oldClose := addRecord, closeStore
|
||||
t.Cleanup(func() { addRecord, closeStore = oldAdd, oldClose })
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,464 @@
|
||||
//go:build linux
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"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 verifiedKeyIDHeader = "X-DWH-Key-ID"
|
||||
|
||||
func TestVerifyDecisionsAreFailClosedAndBodiesEmpty(t *testing.T) {
|
||||
store, closeStore := testStore(t)
|
||||
defer closeStore()
|
||||
valid, validRecord := testV1(t, 1)
|
||||
if err := store.Add(validRecord); err != nil {
|
||||
t.Fatalf("Add(valid) error = %v", err)
|
||||
}
|
||||
legacy := []byte("synthetic-legacy-secret")
|
||||
if err := store.Add(testLegacy(legacy)); err != nil {
|
||||
t.Fatalf("Add(legacy) error = %v", err)
|
||||
}
|
||||
expired, expiredRecord := testV1(t, 2)
|
||||
expiresAt := time.Date(2026, 8, 20, 11, 0, 0, 0, time.UTC)
|
||||
expiredRecord.ExpiresAt = &expiresAt
|
||||
if err := store.Add(expiredRecord); err != nil {
|
||||
t.Fatalf("Add(expired) error = %v", err)
|
||||
}
|
||||
revoked, revokedRecord := testV1(t, 3)
|
||||
if err := store.Add(revokedRecord); err != nil {
|
||||
t.Fatalf("Add(revoked) error = %v", err)
|
||||
}
|
||||
if err := store.Revoke(revokedRecord.KeyID, "synthetic", revokedRecord.CreatedAt.Add(time.Hour)); err != nil {
|
||||
t.Fatalf("Revoke() error = %v", err)
|
||||
}
|
||||
|
||||
var logs bytes.Buffer
|
||||
handler := New(store, log.New(&logs, "", 0), func() time.Time {
|
||||
return time.Date(2026, 8, 20, 12, 0, 0, 0, time.UTC)
|
||||
})
|
||||
cases := []struct {
|
||||
name string
|
||||
headers []string
|
||||
status int
|
||||
keyID string
|
||||
logCode string
|
||||
}{
|
||||
{name: "valid v1", headers: []string{string(valid)}, status: http.StatusNoContent, keyID: validRecord.KeyID, logCode: "allow_v1"},
|
||||
{name: "valid legacy", headers: []string{string(legacy)}, status: http.StatusNoContent, keyID: record.LegacyKeyID, logCode: "allow_legacy"},
|
||||
{name: "absent", status: http.StatusUnauthorized, logCode: "deny_header"},
|
||||
{name: "duplicate", headers: []string{string(valid), string(valid)}, status: http.StatusUnauthorized, logCode: "deny_header"},
|
||||
{name: "malformed", headers: []string{"thtdwh_v1.bad.bad"}, status: http.StatusUnauthorized, logCode: "deny_v1"},
|
||||
{name: "unknown", headers: []string{string(testUnknown(valid))}, status: http.StatusUnauthorized, logCode: "deny_unknown"},
|
||||
{name: "changed", headers: []string{string(testChanged(valid))}, status: http.StatusUnauthorized, logCode: "deny_mismatch"},
|
||||
{name: "expired", headers: []string{string(expired)}, status: http.StatusUnauthorized, logCode: "deny_unknown"},
|
||||
{name: "revoked", headers: []string{string(revoked)}, status: http.StatusUnauthorized, logCode: "deny_revoked"},
|
||||
{name: "oversized", headers: []string{strings.Repeat("a", credential.MaxHeaderBytes+1)}, status: http.StatusUnauthorized, logCode: "deny_header"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
response := verifyRequest(handler, tc.headers...)
|
||||
if response.Code != tc.status {
|
||||
t.Fatalf("status = %d, want %d", response.Code, tc.status)
|
||||
}
|
||||
if response.Body.Len() != 0 {
|
||||
t.Fatalf("body = %q, want empty", response.Body.String())
|
||||
}
|
||||
if got := response.Header().Get(verifiedKeyIDHeader); got != tc.keyID {
|
||||
t.Fatalf("%s = %q, want %q", verifiedKeyIDHeader, got, tc.keyID)
|
||||
}
|
||||
if !strings.Contains(logs.String(), tc.logCode) {
|
||||
t.Fatalf("logs = %q, want decision %q", logs.String(), tc.logCode)
|
||||
}
|
||||
})
|
||||
}
|
||||
for _, forbidden := range []string{string(valid), string(legacy), validRecord.SecretSHA256, validRecord.Description, "query-secret"} {
|
||||
if strings.Contains(logs.String(), forbidden) {
|
||||
t.Fatalf("logs exposed forbidden value %q: %q", forbidden, logs.String())
|
||||
}
|
||||
}
|
||||
for _, line := range strings.Split(strings.TrimSpace(logs.String()), "\n") {
|
||||
fields := strings.Fields(line)
|
||||
if len(fields) != 2 && len(fields) != 3 {
|
||||
t.Fatalf("log fields = %q, want timestamp+decision with optional public ID", line)
|
||||
}
|
||||
if !strings.HasPrefix(fields[0], "timestamp=2026-08-20T12:00:00Z") || !strings.HasPrefix(fields[1], "decision=") {
|
||||
t.Fatalf("log line = %q, want timestamp and decision only", line)
|
||||
}
|
||||
if len(fields) == 3 && !strings.HasPrefix(fields[2], "key_id=") {
|
||||
t.Fatalf("log public field = %q, want key_id", line)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyMapsRegistryIntegrityFailuresToServiceUnavailable(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
store, err := registry.Open(root)
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
defer store.Close()
|
||||
value, item := testV1(t, 4)
|
||||
if err := store.Add(item); err != nil {
|
||||
t.Fatalf("Add() error = %v", err)
|
||||
}
|
||||
if err := os.Chmod(filepath.Join(root, "active", item.KeyID+".json"), 0o660); err != nil {
|
||||
t.Fatalf("Chmod() error = %v", err)
|
||||
}
|
||||
response := verifyRequest(New(store, log.New(&bytes.Buffer{}, "", 0), time.Now), string(value))
|
||||
if response.Code != http.StatusServiceUnavailable || response.Body.Len() != 0 {
|
||||
t.Fatalf("integrity response = (%d, %q), want (503, empty)", response.Code, response.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyFailsClosedWhenUnrelatedRegistryRecordIsCorrupt(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
store, err := registry.Open(root)
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
defer store.Close()
|
||||
value, item := testV1(t, 5)
|
||||
if err := store.Add(item); err != nil {
|
||||
t.Fatalf("Add() error = %v", err)
|
||||
}
|
||||
_, corrupt := testV1(t, 6)
|
||||
path := filepath.Join(root, "active", corrupt.KeyID+".json")
|
||||
if err := os.WriteFile(path, []byte(`{"schema_version":`), 0o640); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
if err := os.Chmod(path, 0o640); err != nil {
|
||||
t.Fatalf("Chmod() error = %v", err)
|
||||
}
|
||||
response := verifyRequest(New(store, log.New(&bytes.Buffer{}, "", 0), time.Now), string(value))
|
||||
if response.Code != http.StatusServiceUnavailable || response.Body.Len() != 0 {
|
||||
t.Fatalf("unrelated corrupt record response = (%d, %q), want (503, empty)", response.Code, response.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyFailsClosedWhenActiveAndRevokedLegacyRecordsCoexist(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
store, err := registry.Open(root)
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
defer store.Close()
|
||||
legacyValue := []byte("synthetic-legacy-secret")
|
||||
legacy := testLegacy(legacyValue)
|
||||
if err := store.Add(legacy); err != nil {
|
||||
t.Fatalf("Add() error = %v", err)
|
||||
}
|
||||
revokedAt := legacy.CreatedAt.Add(time.Hour)
|
||||
legacy.RevokedAt = &revokedAt
|
||||
legacy.RevocationReason = "synthetic"
|
||||
data, err := json.Marshal(legacy)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal() error = %v", err)
|
||||
}
|
||||
data = append(data, '\n')
|
||||
path := filepath.Join(root, "revoked", record.LegacyKeyID+".json")
|
||||
if err := os.WriteFile(path, data, 0o640); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
if err := os.Chmod(path, 0o640); err != nil {
|
||||
t.Fatalf("Chmod() error = %v", err)
|
||||
}
|
||||
response := verifyRequest(New(store, log.New(&bytes.Buffer{}, "", 0), time.Now), string(legacyValue))
|
||||
if response.Code != http.StatusServiceUnavailable || response.Body.Len() != 0 {
|
||||
t.Fatalf("multiple legacy response = (%d, %q), want (503, empty)", response.Code, response.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyMapsMultipleLegacyRecordsToServiceUnavailable(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
store, err := registry.Open(root)
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
defer store.Close()
|
||||
legacy := testLegacy([]byte("synthetic-legacy-secret"))
|
||||
if err := store.Add(legacy); err != nil {
|
||||
t.Fatalf("Add() error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(root, "active", "duplicate.json"), []byte(`{"schema_version":1}`), 0o640); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
if err := os.Chmod(filepath.Join(root, "active", "duplicate.json"), 0o640); err != nil {
|
||||
t.Fatalf("Chmod() error = %v", err)
|
||||
}
|
||||
response := verifyRequest(New(store, log.New(&bytes.Buffer{}, "", 0), time.Now), "synthetic-legacy-secret")
|
||||
if response.Code != http.StatusServiceUnavailable || response.Body.Len() != 0 {
|
||||
t.Fatalf("multiple legacy response = (%d, %q), want (503, empty)", response.Code, response.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestListenAndServeFailsClosedForInvalidPathsAndCollisions(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
parent := t.TempDir()
|
||||
for _, tc := range []struct {
|
||||
name, registryRoot, socket string
|
||||
setup func(string)
|
||||
}{
|
||||
{name: "noncanonical registry root", registryRoot: root + "/.", socket: filepath.Join(parent, "verify.sock")},
|
||||
{name: "noncanonical socket", registryRoot: root, socket: parent + "/./verify.sock"},
|
||||
{name: "regular collision", registryRoot: root, socket: filepath.Join(parent, "regular.sock"), setup: func(path string) {
|
||||
if err := os.WriteFile(path, []byte("do not remove"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}},
|
||||
{name: "directory collision", registryRoot: root, socket: filepath.Join(parent, "directory.sock"), setup: func(path string) {
|
||||
if err := os.Mkdir(path, 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}},
|
||||
{name: "symlink collision", registryRoot: root, socket: filepath.Join(parent, "symlink.sock"), setup: func(path string) {
|
||||
if err := os.Symlink(filepath.Join(parent, "target.sock"), path); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if tc.setup != nil {
|
||||
tc.setup(tc.socket)
|
||||
}
|
||||
err := ListenAndServe(context.Background(), Config{RegistryRoot: tc.registryRoot, Socket: tc.socket, Logger: log.New(&bytes.Buffer{}, "", 0), Now: time.Now})
|
||||
if err == nil {
|
||||
t.Fatal("ListenAndServe() error = nil, want fail-closed refusal")
|
||||
}
|
||||
if tc.setup != nil {
|
||||
if _, statErr := os.Lstat(tc.socket); statErr != nil {
|
||||
t.Fatalf("collision was removed: %v", statErr)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestListenAndServeReclaimsOnlyOwnedStaleSocketAndCleansUpOnCancellation(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
socket := filepath.Join(t.TempDir(), "verify.sock")
|
||||
stale, err := net.ListenUnix("unix", &net.UnixAddr{Name: socket, Net: "unix"})
|
||||
if err != nil {
|
||||
t.Fatalf("ListenUnix() error = %v", err)
|
||||
}
|
||||
stale.SetUnlinkOnClose(false)
|
||||
if err := stale.Close(); err != nil {
|
||||
t.Fatalf("Close(stale) error = %v", err)
|
||||
}
|
||||
if info, err := os.Lstat(socket); err != nil || info.Mode()&os.ModeSocket == 0 {
|
||||
t.Fatalf("stale socket = (%v, %v), want socket", info, err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
errs := make(chan error, 1)
|
||||
go func() {
|
||||
errs <- ListenAndServe(ctx, Config{RegistryRoot: root, Socket: socket, Logger: log.New(&bytes.Buffer{}, "", 0), Now: time.Now})
|
||||
}()
|
||||
waitForReadySocket(t, socket)
|
||||
info, err := os.Stat(socket)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat(socket) error = %v", err)
|
||||
}
|
||||
if got, want := info.Mode().Perm(), os.FileMode(0o660); got != want {
|
||||
t.Fatalf("socket mode = %04o, want %04o", got, want)
|
||||
}
|
||||
cancel()
|
||||
if err := <-errs; err != nil {
|
||||
t.Fatalf("ListenAndServe() error = %v", err)
|
||||
}
|
||||
if _, err := os.Lstat(socket); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("socket after cancellation = %v, want removal", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListenAndServeRefusesStaleSocketWithForeignOwner(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
socket := filepath.Join(t.TempDir(), "verify.sock")
|
||||
stale, err := net.ListenUnix("unix", &net.UnixAddr{Name: socket, Net: "unix"})
|
||||
if err != nil {
|
||||
t.Fatalf("ListenUnix() error = %v", err)
|
||||
}
|
||||
stale.SetUnlinkOnClose(false)
|
||||
if err := stale.Close(); err != nil {
|
||||
t.Fatalf("Close(stale) error = %v", err)
|
||||
}
|
||||
previous := socketOwner
|
||||
socketOwner = func(string) (uint32, error) { return uint32(os.Geteuid()) + 1, nil }
|
||||
defer func() { socketOwner = previous }()
|
||||
err = ListenAndServe(context.Background(), Config{RegistryRoot: root, Socket: socket, Logger: log.New(&bytes.Buffer{}, "", 0), Now: time.Now})
|
||||
if err == nil {
|
||||
t.Fatal("ListenAndServe() error = nil, want foreign-owner refusal")
|
||||
}
|
||||
if _, err := os.Lstat(socket); err != nil {
|
||||
t.Fatalf("foreign socket was removed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListenAndServeRefusesLiveSocket(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
socket := filepath.Join(t.TempDir(), "verify.sock")
|
||||
live, err := net.ListenUnix("unix", &net.UnixAddr{Name: socket, Net: "unix"})
|
||||
if err != nil {
|
||||
t.Fatalf("ListenUnix() error = %v", err)
|
||||
}
|
||||
defer live.Close()
|
||||
err = ListenAndServe(context.Background(), Config{RegistryRoot: root, Socket: socket, Logger: log.New(&bytes.Buffer{}, "", 0), Now: time.Now})
|
||||
if err == nil {
|
||||
t.Fatal("ListenAndServe() error = nil, want live-socket refusal")
|
||||
}
|
||||
if info, statErr := os.Lstat(socket); statErr != nil || info.Mode()&os.ModeSocket == 0 {
|
||||
t.Fatalf("live socket was changed: (%v, %v)", info, statErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListenAndServeCleanupRefusesChangedSocketPath(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
socket := filepath.Join(t.TempDir(), "verify.sock")
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
errs := make(chan error, 1)
|
||||
go func() {
|
||||
errs <- ListenAndServe(ctx, Config{RegistryRoot: root, Socket: socket, Logger: log.New(&bytes.Buffer{}, "", 0), Now: time.Now})
|
||||
}()
|
||||
waitForReadySocket(t, socket)
|
||||
previous := socketOwner
|
||||
swapped := false
|
||||
t.Cleanup(func() { socketOwner = previous })
|
||||
socketOwner = func(path string) (uint32, error) {
|
||||
if !swapped {
|
||||
swapped = true
|
||||
if err := os.Remove(path); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := os.WriteFile(path, []byte("replacement"), 0o600); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
return uint32(os.Geteuid()), nil
|
||||
}
|
||||
cancel()
|
||||
err := <-errs
|
||||
if err == nil {
|
||||
t.Fatal("ListenAndServe() error = nil, want changed-path refusal")
|
||||
}
|
||||
value, readErr := os.ReadFile(socket)
|
||||
if readErr != nil || string(value) != "replacement" {
|
||||
t.Fatalf("replacement after cleanup = (%q, %v), want retained regular file", value, readErr)
|
||||
}
|
||||
}
|
||||
|
||||
func waitForReadySocket(t *testing.T, socket string) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if info, err := os.Lstat(socket); err == nil && info.Mode()&os.ModeSocket != 0 && info.Mode().Perm() == 0o660 {
|
||||
return
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("socket %q did not become ready with mode 0660", socket)
|
||||
}
|
||||
|
||||
func TestVerifyRejectsOtherPathsAndMethodsWithEmptyBodies(t *testing.T) {
|
||||
store, closeStore := testStore(t)
|
||||
defer closeStore()
|
||||
handler := New(store, log.New(&bytes.Buffer{}, "", 0), time.Now)
|
||||
for _, tc := range []struct {
|
||||
method string
|
||||
path string
|
||||
status int
|
||||
}{
|
||||
{method: http.MethodGet, path: "/other", status: http.StatusNotFound},
|
||||
{method: http.MethodPost, path: "/verify", status: http.StatusMethodNotAllowed},
|
||||
} {
|
||||
t.Run(tc.method+tc.path, func(t *testing.T) {
|
||||
request := httptest.NewRequest(tc.method, tc.path, nil)
|
||||
response := httptest.NewRecorder()
|
||||
handler.ServeHTTP(response, request)
|
||||
if response.Code != tc.status || response.Body.Len() != 0 {
|
||||
t.Fatalf("response = (%d, %q), want (%d, empty)", response.Code, response.Body.String(), tc.status)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func verifyRequest(handler http.Handler, values ...string) *httptest.ResponseRecorder {
|
||||
request := httptest.NewRequest(http.MethodGet, "/verify?credential=query-secret", nil)
|
||||
for _, value := range values {
|
||||
request.Header.Add("X-API-Key", value)
|
||||
}
|
||||
response := httptest.NewRecorder()
|
||||
handler.ServeHTTP(response, request)
|
||||
return response
|
||||
}
|
||||
|
||||
func testStore(t *testing.T) (*registry.Store, func()) {
|
||||
t.Helper()
|
||||
store, err := registry.Open(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
return store, func() {
|
||||
if err := store.Close(); err != nil {
|
||||
t.Fatalf("Close() error = %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func testV1(t *testing.T, seed byte) ([]byte, record.Record) {
|
||||
t.Helper()
|
||||
material, err := credential.Generate(bytes.NewReader(bytes.Repeat([]byte{seed}, 44)))
|
||||
if err != nil {
|
||||
t.Fatalf("Generate() error = %v", err)
|
||||
}
|
||||
return material.Value, record.Record{
|
||||
SchemaVersion: record.SchemaVersion,
|
||||
Kind: record.KindV1,
|
||||
KeyID: material.KeyID,
|
||||
InstallationID: "test-installation",
|
||||
Description: "synthetic description",
|
||||
SecretSHA256: base64.RawURLEncoding.EncodeToString(material.Digest[:]),
|
||||
CreatedAt: time.Date(2026, 8, 20, 10, 0, int(seed), 0, time.UTC),
|
||||
}
|
||||
}
|
||||
|
||||
func testLegacy(value []byte) record.Record {
|
||||
digest := sha256.Sum256(value)
|
||||
return record.Record{
|
||||
SchemaVersion: record.SchemaVersion,
|
||||
Kind: record.KindLegacyRaw,
|
||||
KeyID: record.LegacyKeyID,
|
||||
InstallationID: record.LegacyKeyID,
|
||||
Description: "legacy synthetic description",
|
||||
SecretSHA256: base64.RawURLEncoding.EncodeToString(digest[:]),
|
||||
CreatedAt: time.Date(2026, 8, 20, 10, 0, 0, 0, time.UTC),
|
||||
}
|
||||
}
|
||||
|
||||
func testChanged(value []byte) []byte {
|
||||
changed := append([]byte(nil), value...)
|
||||
changed[len(credential.Prefix)+1+credential.KeyIDEncodedLength+1] = 'B'
|
||||
return changed
|
||||
}
|
||||
|
||||
func testUnknown(value []byte) []byte {
|
||||
unknown := append([]byte(nil), value...)
|
||||
unknown[len(credential.Prefix)+1] = 'B'
|
||||
return unknown
|
||||
}
|
||||
Reference in New Issue
Block a user