//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() initializeRegistry(t, root) 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 TestListenAndServeRejectsUnsafeSocketParents(t *testing.T) { root := t.TempDir() initializeRegistry(t, root) for _, mode := range []os.FileMode{0o770, 0o702} { t.Run("mode "+mode.String(), func(t *testing.T) { parent := t.TempDir() if err := os.Chmod(parent, mode); err != nil { t.Fatalf("Chmod(%v) error = %v", mode, err) } socket := filepath.Join(parent, "verify.sock") err := ListenAndServe(context.Background(), Config{RegistryRoot: root, Socket: socket, Logger: log.New(&bytes.Buffer{}, "", 0), Now: time.Now}) if err == nil { t.Fatalf("ListenAndServe(parent mode %v) error = nil, want refusal", mode) } if _, err := os.Lstat(socket); !errors.Is(err, os.ErrNotExist) { t.Fatalf("unsafe parent created socket: %v", err) } }) } } func TestListenAndServeRejectsForeignOwnedSocketParent(t *testing.T) { root := t.TempDir() initializeRegistry(t, root) parent := t.TempDir() socket := filepath.Join(parent, "verify.sock") previous := socketParentOwner t.Cleanup(func() { socketParentOwner = previous }) socketParentOwner = func(string) (uint32, error) { return uint32(os.Geteuid()) + 1, nil } 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-parent refusal") } if _, err := os.Lstat(socket); !errors.Is(err, os.ErrNotExist) { t.Fatalf("foreign parent created socket: %v", err) } } func TestListenAndServeReclaimsOnlyOwnedStaleSocketAndCleansUpOnCancellation(t *testing.T) { root := t.TempDir() initializeRegistry(t, root) 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() initializeRegistry(t, root) 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() initializeRegistry(t, root) 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() initializeRegistry(t, root) 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 initializeRegistry(t *testing.T, root string) { t.Helper() store, err := registry.Open(root) if err != nil { t.Fatalf("Open() error = %v", err) } if err := store.Close(); err != nil { t.Fatalf("Close() error = %v", err) } } 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 }