feat: serve DWH authentication over Unix socket

This commit is contained in:
User
2026-08-21 01:57:19 +02:00
parent 943f809d01
commit 419c3440d7
4 changed files with 855 additions and 16 deletions
+19 -13
View File
@@ -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 })
+344
View File
@@ -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
}