diff --git a/tools/dwh-auth/internal/command/command.go b/tools/dwh-auth/internal/command/command.go index 6d24bd52..102b6c6e 100644 --- a/tools/dwh-auth/internal/command/command.go +++ b/tools/dwh-auth/internal/command/command.go @@ -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 { diff --git a/tools/dwh-auth/internal/command/command_test.go b/tools/dwh-auth/internal/command/command_test.go index 1059c1c6..f3174d31 100644 --- a/tools/dwh-auth/internal/command/command_test.go +++ b/tools/dwh-auth/internal/command/command_test.go @@ -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 }) diff --git a/tools/dwh-auth/internal/service/service.go b/tools/dwh-auth/internal/service/service.go new file mode 100644 index 00000000..73b1b9f7 --- /dev/null +++ b/tools/dwh-auth/internal/service/service.go @@ -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) +} diff --git a/tools/dwh-auth/internal/service/service_test.go b/tools/dwh-auth/internal/service/service_test.go new file mode 100644 index 00000000..f963a415 --- /dev/null +++ b/tools/dwh-auth/internal/service/service_test.go @@ -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 +}