Files
ThothII/tools/dwh-auth/internal/service/service_test.go

519 lines
18 KiB
Go

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