From 7e61c31b2da9c2bd2c5f7e6f870b6907cc631326 Mon Sep 17 00:00:00 2001 From: peg Date: Thu, 20 Aug 2026 09:47:50 +0200 Subject: [PATCH 1/9] Support image hash measurements --- Makefile | 2 +- README.md | 29 +++++ adapters/database/service.go | 33 ++++-- adapters/database/service_test.go | 31 +++++ adapters/database/types.go | 32 ++++-- application/service.go | 163 ++++++++++++++++++++++++++- application/service_test.go | 126 ++++++++++++++++++++- docker/mock-proxy/nginx-default.conf | 4 +- docs/devenv-setup.md | 2 +- domain/types.go | 55 +++++++++ httpserver/e2e_test.go | 73 ++++++++++-- ports/admin_handler.go | 7 +- ports/http_handler.go | 103 ++++++++++++++++- ports/types.go | 31 ++++- ports/types_test.go | 110 ++++++++++++++++++ schema/005_dcap_image_hashes.sql | 6 + scripts/ci/e2e-test.hurl | 4 +- testdata/get-measurements.json | 11 ++ 18 files changed, 771 insertions(+), 51 deletions(-) create mode 100644 schema/005_dcap_image_hashes.sql diff --git a/Makefile b/Makefile index fcdc98a..29902f6 100644 --- a/Makefile +++ b/Makefile @@ -118,7 +118,7 @@ db-dump: ## Dump the database contents to file 'database.dump' .PHONY: dev-db-setup dev-db-setup: ## Create the basic database entries for testing and development @printf "$(BLUE)Create the allow-all measurements $(NC)\n" - $(CURL) $(CURL_AUTH) --request POST --url http://localhost:8081/api/admin/v1/measurements --data '{"measurement_id": "test1","attestation_type": "test","measurements": {}}' + $(CURL) $(CURL_AUTH) --request POST --url http://localhost:8081/api/admin/v1/measurements --data '{"measurement_id": "test1","attestation_type": "dcap-tdx","measurements": {}}' @printf "$(BLUE)Enable the measurements $(NC)\n" $(CURL) $(CURL_AUTH) --request POST --url http://localhost:8081/api/admin/v1/measurements/activation/test1 --data '{"enabled": true}' diff --git a/README.md b/README.md index 0df70a3..8cd0d04 100644 --- a/README.md +++ b/README.md @@ -207,6 +207,19 @@ Response: Array with currently allowed measurement JSONs [testdata/get-measurements.json](https://github.com/flashbots/builder-config-hub/blob/main/testdata/get-measurements.json) +Authenticated instance requests receive two headers from the attestation proxy: + +- `X-Flashbots-Attestation-Type`, containing `dcap-tdx`, `gcp-tdx`, `azure-tdx`, or `none`. +- `X-Flashbots-Measurement`, containing the expected policy selected by the proxy. + +For example, a DCAP policy is represented as compact JSON with an array of accepted values per register: + +```json +{"type":"dcap","measurements":{"0":[""],"3":[""]}} +``` + +Portable image policies use `type: "image"`; `no_attestation` has no `measurements` field. Legacy headers containing a plain register-to-string object remain accepted during migration. + --- ## Admin Endpoints @@ -253,6 +266,22 @@ To allow _any_ measurement, use an empty measurements field: } ``` +Portable DCAP policies use image-component hashes instead of register values. They are supported for both `dcap-tdx` and `gcp-tdx`; `dcap_image_hashes` and `measurements` cannot be used together. + +```json +{ + "measurement_id": "portable-image-v1", + "attestation_type": "dcap-tdx", + "dcap_image_hashes": { + "uki_authenticode": "111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111", + "kernel_authenticode": "222222222222222222222222222222222222222222222222222222222222222222222222222222222222222222222222", + "cmdline_hash": "333333333333333333333333333333333333333333333333333333333333333333333333333333333333333333333333", + "initrd_hash": "444444444444444444444444444444444444444444444444444444444444444444444444444444444444444444444444", + "gpt_disk_guid_hash": "555555555555555555555555555555555555555555555555555555555555555555555555555555555555555555555555" + } +} +``` + ### Enable/disable measurements `POST /api/admin/v1/measurements/activation/{measurement_id}` diff --git a/adapters/database/service.go b/adapters/database/service.go index 5e11472..3377dd5 100644 --- a/adapters/database/service.go +++ b/adapters/database/service.go @@ -39,7 +39,7 @@ func (s *Service) Close() error { func (s *Service) GetActiveMeasurementsByType(ctx context.Context, attestationType string) ([]domain.Measurement, error) { var measurements []Measurement - err := s.DB.SelectContext(ctx, &measurements, `SELECT * FROM measurements_whitelist WHERE is_active=true AND attestation_type=$1`, attestationType) + err := s.DB.SelectContext(ctx, &measurements, `SELECT * FROM measurements_whitelist WHERE is_active=true AND attestation_type=$1 ORDER BY id`, attestationType) var domainMeasurements []domain.Measurement for _, m := range measurements { domainM, err := convertMeasurementToDomain(m) @@ -76,7 +76,7 @@ func (s *Service) GetBuilderByIP(ip net.IP) (*domain.Builder, error) { // GetActiveMeasurements retrieves all measurements func (s *Service) GetActiveMeasurements(ctx context.Context) ([]domain.Measurement, error) { var measurements []Measurement - err := s.DB.SelectContext(ctx, &measurements, `SELECT * FROM measurements_whitelist WHERE is_active=true`) + err := s.DB.SelectContext(ctx, &measurements, `SELECT * FROM measurements_whitelist WHERE is_active=true ORDER BY id`) var domainMeasurements []domain.Measurement for _, m := range measurements { domainM, err := convertMeasurementToDomain(m) @@ -245,14 +245,29 @@ func (s *Service) LogEvent(ctx context.Context, eventName, builderName, name str } func (s *Service) AddMeasurement(ctx context.Context, measurement domain.Measurement, enabled bool) error { - bts, err := json.Marshal(measurement.Measurement) - if err != nil { - return err + var measurementJSON any + var imageHashesJSON any + if measurement.DcapImageHashes != nil { + bts, err := json.Marshal(measurement.DcapImageHashes) + if err != nil { + return err + } + imageHashesJSON = bts + } else { + measurements := measurement.Measurement + if measurements == nil { + measurements = make(map[string]domain.SingleMeasurement) + } + bts, err := json.Marshal(measurements) + if err != nil { + return err + } + measurementJSON = bts } - _, err = s.DB.ExecContext(ctx, ` - INSERT INTO measurements_whitelist (name, attestation_type, measurement, is_active) - VALUES ($1, $2, $3, $4) - `, measurement.Name, measurement.AttestationType, bts, enabled) + _, err := s.DB.ExecContext(ctx, ` + INSERT INTO measurements_whitelist (name, attestation_type, measurement, dcap_image_hashes, is_active) + VALUES ($1, $2, $3, $4, $5) + `, measurement.Name, measurement.AttestationType, measurementJSON, imageHashesJSON, enabled) return err } diff --git a/adapters/database/service_test.go b/adapters/database/service_test.go index 1432c43..99df124 100644 --- a/adapters/database/service_test.go +++ b/adapters/database/service_test.go @@ -5,6 +5,7 @@ import ( "encoding/json" "net" "os" + "strings" "testing" "github.com/flashbots/builder-hub/domain" @@ -48,6 +49,36 @@ func TestGetBuilder(t *testing.T) { }) } +func TestPortableMeasurement(t *testing.T) { + if os.Getenv("RUN_DB_TESTS") != "1" { + t.Skip("skipping test; RUN_DB_TESTS is not set to 1") + } + dbService, err := NewDatabaseService("postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable") + require.NoError(t, err) + _, err = dbService.DB.Exec("TRUNCATE TABLE public.measurements_whitelist CASCADE") + require.NoError(t, err) + + hashes := &domain.DcapImageHashes{ + UKIAuthenticode: strings.Repeat("1", 96), + KernelAuthenticode: strings.Repeat("2", 96), + CmdlineHash: strings.Repeat("3", 96), + InitrdHash: strings.Repeat("4", 96), + GPTDiskGUIDHash: strings.Repeat("5", 96), + } + err = dbService.AddMeasurement(context.Background(), domain.Measurement{ + Name: "portable-image", + AttestationType: "dcap-tdx", + DcapImageHashes: hashes, + }, true) + require.NoError(t, err) + + measurements, err := dbService.GetActiveMeasurementsByType(context.Background(), "dcap-tdx") + require.NoError(t, err) + require.Len(t, measurements, 1) + require.Nil(t, measurements[0].Measurement) + require.Equal(t, hashes, measurements[0].DcapImageHashes) +} + func TestAdminFlow(t *testing.T) { if os.Getenv("RUN_DB_TESTS") != "1" { t.Skip("skipping test; RUN_DB_TESTS is not set to 1") diff --git a/adapters/database/types.go b/adapters/database/types.go index 3aacf7d..86e4f7e 100644 --- a/adapters/database/types.go +++ b/adapters/database/types.go @@ -10,23 +10,31 @@ import ( ) type Measurement struct { - ID int `db:"id"` - Name string `db:"name"` - AttestationType string `db:"attestation_type"` - Measurement json.RawMessage `db:"measurement"` - IsActive bool `db:"is_active"` - CreatedAt time.Time `db:"created_at"` - UpdatedAt time.Time `db:"updated_at"` - DeprecatedAt *time.Time `db:"deprecated_at"` + ID int `db:"id"` + Name string `db:"name"` + AttestationType string `db:"attestation_type"` + Measurement sql.NullString `db:"measurement"` + DcapImageHashes sql.NullString `db:"dcap_image_hashes"` + IsActive bool `db:"is_active"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` + DeprecatedAt *time.Time `db:"deprecated_at"` } func convertMeasurementToDomain(measurement Measurement) (*domain.Measurement, error) { var m domain.Measurement m.AttestationType = measurement.AttestationType - m.Measurement = make(map[string]domain.SingleMeasurement) - err := json.Unmarshal(measurement.Measurement, &m.Measurement) - if err != nil { - return nil, err + if measurement.Measurement.Valid && measurement.Measurement.String != "null" { + m.Measurement = make(map[string]domain.SingleMeasurement) + if err := json.Unmarshal([]byte(measurement.Measurement.String), &m.Measurement); err != nil { + return nil, err + } + } + if measurement.DcapImageHashes.Valid && measurement.DcapImageHashes.String != "null" { + m.DcapImageHashes = new(domain.DcapImageHashes) + if err := json.Unmarshal([]byte(measurement.DcapImageHashes.String), m.DcapImageHashes); err != nil { + return nil, err + } } m.Name = measurement.Name return &m, nil diff --git a/application/service.go b/application/service.go index 2972171..c3384bc 100644 --- a/application/service.go +++ b/application/service.go @@ -6,6 +6,8 @@ import ( "errors" "fmt" "net" + "strconv" + "strings" "github.com/flashbots/builder-hub/domain" ) @@ -63,7 +65,7 @@ func (b *BuilderHub) GetConfigWithSecrets(ctx context.Context, builderName strin return secr, nil } -func (b *BuilderHub) VerifyIPAndMeasurements(ctx context.Context, ip net.IP, measurement map[string]string, attestationType string) (*domain.Builder, string, error) { +func (b *BuilderHub) VerifyIPAndMeasurements(ctx context.Context, ip net.IP, measurement domain.SuppliedMeasurements, attestationType string) (*domain.Builder, string, error) { measurements, err := b.dataAccessor.GetActiveMeasurementsByType(ctx, attestationType) if err != nil { return nil, "", fmt.Errorf("failing to fetch corresponding measurement data %s %w", attestationType, err) @@ -81,15 +83,170 @@ func (b *BuilderHub) VerifyIPAndMeasurements(ctx context.Context, ip net.IP, mea return builder, measurementName, nil } -func validateMeasurement(measurement map[string]string, measurementTemplate []domain.Measurement) (string, error) { +func validateMeasurement(measurement domain.SuppliedMeasurements, measurementTemplate []domain.Measurement) (string, error) { for _, m := range measurementTemplate { - if checkMeasurement(measurement, m) { + matched := false + if measurement.Type == domain.ExpectedMeasurementLegacy { + matched = checkMeasurement(measurement.Legacy, m) + } else { + matched = checkExpectedMeasurement(measurement, m) + } + if matched { return m.Name, nil } } return "", domain.ErrNotFound } +func checkExpectedMeasurement(measurement domain.SuppliedMeasurements, template domain.Measurement) bool { + expectedType, ok := expectedMeasurementType(template) + if !ok || measurement.Type != expectedType { + return false + } + + switch expectedType { + case domain.ExpectedMeasurementDCAP, domain.ExpectedMeasurementAzure: + supplied, ok := normalizeSuppliedRegisters(measurement.Registers, template.AttestationType) + if !ok { + return false + } + expected, ok := normalizeExpectedRegisters(template.Measurement, template.AttestationType) + return ok && equalRegisterSets(supplied, expected) + case domain.ExpectedMeasurementImage: + return equalImageHashes(measurement.DcapImageHashes, template.DcapImageHashes) + case domain.ExpectedMeasurementNoAttestation: + return true + default: + return false + } +} + +func expectedMeasurementType(template domain.Measurement) (domain.ExpectedMeasurementType, bool) { + if template.DcapImageHashes != nil { + if template.AttestationType == "dcap-tdx" || template.AttestationType == "gcp-tdx" { + return domain.ExpectedMeasurementImage, true + } + return "", false + } + + switch template.AttestationType { + case "dcap-tdx", "gcp-tdx": + return domain.ExpectedMeasurementDCAP, true + case "azure-tdx": + return domain.ExpectedMeasurementAzure, true + case "none": + return domain.ExpectedMeasurementNoAttestation, true + default: + return "", false + } +} + +func normalizeSuppliedRegisters(registers map[string][]string, attestationType string) (map[string]map[string]struct{}, bool) { + normalized := make(map[string]map[string]struct{}, len(registers)) + for key, values := range registers { + canonical, ok := canonicalRegisterKey(key, attestationType) + if !ok || len(values) == 0 { + return nil, false + } + if _, duplicate := normalized[canonical]; duplicate { + return nil, false + } + normalized[canonical] = stringSet(values) + } + return normalized, true +} + +func normalizeExpectedRegisters(registers map[string]domain.SingleMeasurement, attestationType string) (map[string]map[string]struct{}, bool) { + normalized := make(map[string]map[string]struct{}, len(registers)) + for key, measurement := range registers { + canonical, ok := canonicalRegisterKey(key, attestationType) + values := measurement.GetExpectedValues() + if !ok || len(values) == 0 { + return nil, false + } + if _, duplicate := normalized[canonical]; duplicate { + return nil, false + } + normalized[canonical] = stringSet(values) + } + return normalized, true +} + +func canonicalRegisterKey(key, attestationType string) (string, bool) { + switch attestationType { + case "dcap-tdx", "gcp-tdx": + if index, err := strconv.Atoi(key); err == nil { + if index >= 0 && index <= 4 { + return strconv.Itoa(index), true + } + return "", false + } + switch strings.ToLower(key) { + case "mrtd": + return "0", true + case "rtmr0": + return "1", true + case "rtmr1": + return "2", true + case "rtmr2": + return "3", true + case "rtmr3": + return "4", true + default: + return "", false + } + case "azure-tdx": + indexString := key + if len(key) >= 3 && strings.EqualFold(key[:3], "pcr") { + indexString = key[3:] + } + index, err := strconv.Atoi(indexString) + if err != nil || index < 0 || index > 23 { + return "", false + } + return strconv.Itoa(index), true + default: + return "", false + } +} + +func stringSet(values []string) map[string]struct{} { + set := make(map[string]struct{}, len(values)) + for _, value := range values { + set[strings.ToLower(value)] = struct{}{} + } + return set +} + +func equalRegisterSets(left, right map[string]map[string]struct{}) bool { + if len(left) != len(right) { + return false + } + for key, leftValues := range left { + rightValues, ok := right[key] + if !ok || len(leftValues) != len(rightValues) { + return false + } + for value := range leftValues { + if _, ok := rightValues[value]; !ok { + return false + } + } + } + return true +} + +func equalImageHashes(left, right *domain.DcapImageHashes) bool { + if left == nil || right == nil { + return false + } + return strings.EqualFold(left.UKIAuthenticode, right.UKIAuthenticode) && + strings.EqualFold(left.KernelAuthenticode, right.KernelAuthenticode) && + strings.EqualFold(left.CmdlineHash, right.CmdlineHash) && + strings.EqualFold(left.InitrdHash, right.InitrdHash) && + strings.EqualFold(left.GPTDiskGUIDHash, right.GPTDiskGUIDHash) +} + // validates that all fields from measurementTemplate are the same in measurement. // For each field, the measurement value must match at least one of the expected values (OR semantics). func checkMeasurement(measurement map[string]string, measurementTemplate domain.Measurement) bool { diff --git a/application/service_test.go b/application/service_test.go index 2af59ca..912a45e 100644 --- a/application/service_test.go +++ b/application/service_test.go @@ -1,6 +1,7 @@ package application import ( + "strings" "testing" "github.com/flashbots/builder-hub/domain" @@ -126,7 +127,7 @@ func TestValidateMeasurement(t *testing.T) { "8": "0000", "11": "aaaa", } - name, err := validateMeasurement(measurement, templates) + name, err := validateMeasurement(legacyMeasurements(measurement), templates) require.NoError(t, err) require.Equal(t, "template-1", name) }) @@ -136,7 +137,7 @@ func TestValidateMeasurement(t *testing.T) { "8": "0000", "11": "bbbb", } - name, err := validateMeasurement(measurement, templates) + name, err := validateMeasurement(legacyMeasurements(measurement), templates) require.NoError(t, err) require.Equal(t, "template-1", name) }) @@ -146,7 +147,7 @@ func TestValidateMeasurement(t *testing.T) { "8": "1111", "11": "cccc", } - name, err := validateMeasurement(measurement, templates) + name, err := validateMeasurement(legacyMeasurements(measurement), templates) require.NoError(t, err) require.Equal(t, "template-2", name) }) @@ -156,11 +157,128 @@ func TestValidateMeasurement(t *testing.T) { "8": "9999", "11": "zzzz", } - _, err := validateMeasurement(measurement, templates) + _, err := validateMeasurement(legacyMeasurements(measurement), templates) require.ErrorIs(t, err, domain.ErrNotFound) }) } +func TestCheckExpectedMeasurementRegisters(t *testing.T) { + template := domain.Measurement{ + Name: "azure-policy", + AttestationType: "azure-tdx", + Measurement: map[string]domain.SingleMeasurement{ + "pcr4": {ExpectedAny: []string{"AAAA", "bbbb"}}, + "PCR11": {Expected: "CCCC"}, + }, + } + + t.Run("matches exact normalized policy", func(t *testing.T) { + supplied := domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementAzure, + Registers: map[string][]string{ + "4": {"BBBB", "aaaa"}, + "11": {"cccc"}, + }, + } + require.True(t, checkExpectedMeasurement(supplied, template)) + }) + + t.Run("rejects missing alternative", func(t *testing.T) { + supplied := domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementAzure, + Registers: map[string][]string{ + "4": {"aaaa"}, + "11": {"cccc"}, + }, + } + require.False(t, checkExpectedMeasurement(supplied, template)) + }) + + t.Run("rejects additional register", func(t *testing.T) { + supplied := domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementAzure, + Registers: map[string][]string{ + "4": {"aaaa", "bbbb"}, + "11": {"cccc"}, + "12": {"dddd"}, + }, + } + require.False(t, checkExpectedMeasurement(supplied, template)) + }) + + t.Run("matches DCAP aliases", func(t *testing.T) { + dcapTemplate := domain.Measurement{ + AttestationType: "dcap-tdx", + Measurement: map[string]domain.SingleMeasurement{ + "MRTD": {Expected: "aaaa"}, + "rtmr2": {ExpectedAny: []string{"bbbb", "cccc"}}, + }, + } + supplied := domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementDCAP, + Registers: map[string][]string{ + "0": {"AAAA"}, + "3": {"cccc", "bbbb"}, + }, + } + require.True(t, checkExpectedMeasurement(supplied, dcapTemplate)) + }) +} + +func TestCheckExpectedMeasurementImage(t *testing.T) { + hashes := testImageHashes("a") + supplied := domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementImage, + DcapImageHashes: &hashes, + } + + for _, attestationType := range []string{"dcap-tdx", "gcp-tdx"} { + t.Run(attestationType, func(t *testing.T) { + templateHashes := testImageHashes("A") + template := domain.Measurement{ + AttestationType: attestationType, + DcapImageHashes: &templateHashes, + } + require.True(t, checkExpectedMeasurement(supplied, template)) + + template.DcapImageHashes.InitrdHash = strings.Repeat("b", 96) + require.False(t, checkExpectedMeasurement(supplied, template)) + }) + } +} + +func TestCheckExpectedMeasurementAllowAnyAndNoAttestation(t *testing.T) { + require.True(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementDCAP, + Registers: map[string][]string{}, + }, domain.Measurement{ + AttestationType: "gcp-tdx", + Measurement: map[string]domain.SingleMeasurement{}, + })) + + require.True(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementNoAttestation, + }, domain.Measurement{AttestationType: "none"})) +} + +func legacyMeasurements(values map[string]string) domain.SuppliedMeasurements { + return domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementLegacy, + Legacy: values, + } +} + +func testImageHashes(character string) domain.DcapImageHashes { + value := strings.Repeat(character, 96) + return domain.DcapImageHashes{ + UKIAuthenticode: value, + KernelAuthenticode: value, + CmdlineHash: value, + InitrdHash: value, + GPTDiskGUIDHash: value, + } +} + func TestMatchesAnyExpected(t *testing.T) { t.Run("empty list returns false", func(t *testing.T) { require.False(t, matchesAnyExpected("value", []string{})) diff --git a/docker/mock-proxy/nginx-default.conf b/docker/mock-proxy/nginx-default.conf index 321534b..33c251f 100644 --- a/docker/mock-proxy/nginx-default.conf +++ b/docker/mock-proxy/nginx-default.conf @@ -4,8 +4,8 @@ server { location / { proxy_pass http://builder-hub-api:8080; - proxy_set_header X-Flashbots-Attestation-Type 'test'; - proxy_set_header X-Flashbots-Measurement '{}'; + proxy_set_header X-Flashbots-Attestation-Type 'dcap-tdx'; + proxy_set_header X-Flashbots-Measurement '{"type":"dcap","measurements":{}}'; proxy_set_header X-Forwarded-For '1.2.3.4'; } } diff --git a/docs/devenv-setup.md b/docs/devenv-setup.md index cc3525a..811cee9 100644 --- a/docs/devenv-setup.md +++ b/docs/devenv-setup.md @@ -49,7 +49,7 @@ curl -v \ --url http://localhost:8081/api/admin/v1/measurements \ --data '{ "measurement_id": "test1", - "attestation_type": "test", + "attestation_type": "dcap-tdx", "measurements": {} }' diff --git a/domain/types.go b/domain/types.go index 8970623..705266d 100644 --- a/domain/types.go +++ b/domain/types.go @@ -2,7 +2,9 @@ package domain import ( + "encoding/hex" "errors" + "fmt" "net" "github.com/ethereum/go-ethereum/common" @@ -22,6 +24,58 @@ type Measurement struct { Name string AttestationType string Measurement map[string]SingleMeasurement + DcapImageHashes *DcapImageHashes +} + +// DcapImageHashes contains the image-specific SHA-384 hashes used by portable +// DCAP measurement policies. +type DcapImageHashes struct { + UKIAuthenticode string `json:"uki_authenticode"` + KernelAuthenticode string `json:"kernel_authenticode"` + CmdlineHash string `json:"cmdline_hash"` + InitrdHash string `json:"initrd_hash"` + GPTDiskGUIDHash string `json:"gpt_disk_guid_hash"` +} + +type ExpectedMeasurementType string + +const ( + ExpectedMeasurementLegacy ExpectedMeasurementType = "legacy" + ExpectedMeasurementDCAP ExpectedMeasurementType = "dcap" + ExpectedMeasurementAzure ExpectedMeasurementType = "azure" + ExpectedMeasurementImage ExpectedMeasurementType = "image" + ExpectedMeasurementNoAttestation ExpectedMeasurementType = "no_attestation" +) + +// SuppliedMeasurements represents either a legacy set of observed register +// values or the expected policy selected by a current attestation proxy. +type SuppliedMeasurements struct { + Type ExpectedMeasurementType + Registers map[string][]string + DcapImageHashes *DcapImageHashes + Legacy map[string]string +} + +// Validate checks that every portable image measurement is a SHA-384 value. +func (d DcapImageHashes) Validate() error { + values := []struct { + name string + value string + }{ + {name: "uki_authenticode", value: d.UKIAuthenticode}, + {name: "kernel_authenticode", value: d.KernelAuthenticode}, + {name: "cmdline_hash", value: d.CmdlineHash}, + {name: "initrd_hash", value: d.InitrdHash}, + {name: "gpt_disk_guid_hash", value: d.GPTDiskGUIDHash}, + } + + for _, item := range values { + decoded, err := hex.DecodeString(item.value) + if err != nil || len(decoded) != 48 { + return fmt.Errorf("%s must be a 48-byte hexadecimal value", item.name) + } + } + return nil } // SingleMeasurement represents a single measurement with one or more expected values. @@ -48,6 +102,7 @@ func NewMeasurement(name, attestationType string, measurements map[string]Single AttestationType: attestationType, Measurement: measurements, Name: name, + DcapImageHashes: nil, } } diff --git a/httpserver/e2e_test.go b/httpserver/e2e_test.go index 541c1ef..bb3df31 100644 --- a/httpserver/e2e_test.go +++ b/httpserver/e2e_test.go @@ -9,6 +9,7 @@ import ( "net/http/httptest" "os" "slices" + "strings" "testing" "github.com/ethereum/go-ethereum/common" @@ -47,7 +48,7 @@ func TestCreateMultipleNetworkBuilders(t *testing.T) { measurement := ports.Measurement{ Name: "test-measurement-1", - AttestationType: "test-attestation-type-1", + AttestationType: "azure-tdx", Measurements: map[string]domain.SingleMeasurement{ "8": { Expected: "0000000000000000000000000000000000000000000000000000000000000000", @@ -111,6 +112,20 @@ func TestCreateMultipleNetworkBuilders(t *testing.T) { } }) + t.Run("Auth active builders with expected policy header", func(t *testing.T) { + resp := make([]ports.BuilderWithServiceCreds, 0) + header := map[string]any{ + "type": "azure", + "measurements": map[string][]string{ + "8": {"0000000000000000000000000000000000000000000000000000000000000000"}, + "11": {"efa43e0beff151b0f251c4abf48152382b1452b4414dbd737b4127de05ca31f7"}, + }, + } + status, _ := execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/builders", nil, &resp, measurement.AttestationType, header, "127.0.0.1") + require.Equal(t, http.StatusOK, status) + require.Len(t, resp, 3) + }) + t.Run("Auth active builders other network", func(t *testing.T) { resp := make([]ports.BuilderWithServiceCreds, 0) status, _ := execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/builders", nil, &resp, measurement.AttestationType, map[string]string{"8": "0000000000000000000000000000000000000000000000000000000000000000", "11": "efa43e0beff151b0f251c4abf48152382b1452b4414dbd737b4127de05ca31f7"}, "127.0.1.1") @@ -169,7 +184,7 @@ func TestAuthInteractionFlow(t *testing.T) { createBuilder(t, s, builderName, ip, domain.ProductionNetwork) measurement := ports.Measurement{ Name: "test-measurement-1", - AttestationType: "test-attestation-type-1", + AttestationType: "azure-tdx", Measurements: map[string]domain.SingleMeasurement{ "8": { Expected: "0000000000000000000000000000000000000000000000000000000000000000", @@ -180,10 +195,17 @@ func TestAuthInteractionFlow(t *testing.T) { }, } createMeasurement(t, s, measurement) + expectedHeader := map[string]any{ + "type": "azure", + "measurements": map[string][]string{ + "8": {"0000000000000000000000000000000000000000000000000000000000000000"}, + "11": {"efa43e0beff151b0f251c4abf48152382b1452b4414dbd737b4127de05ca31f7"}, + }, + } t.Run("GetConfig", func(t *testing.T) { resp := make(map[string]string) - status, _ := execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/configuration", nil, &resp, measurement.AttestationType, map[string]string{"8": "0000000000000000000000000000000000000000000000000000000000000000", "11": "efa43e0beff151b0f251c4abf48152382b1452b4414dbd737b4127de05ca31f7"}, "127.0.0.1") + status, _ := execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/configuration", nil, &resp, measurement.AttestationType, expectedHeader, "127.0.0.1") require.Equal(t, http.StatusOK, status) require.Equal(t, builderName+"_test_value_1", resp["test_key_1"]) require.Equal(t, builderName+"_test_secret_value", resp["test_secret_1"]) @@ -210,12 +232,12 @@ func TestAuthInteractionFlow(t *testing.T) { ECDSAPubkey: &addr, Region: "us-east-1", } - status, _ := execRequestAuth(t, s.GetRouter(), http.MethodPost, "/api/l1-builder/v1/register_credentials/rbuilder", sc, nil, measurement.AttestationType, map[string]string{"8": "0000000000000000000000000000000000000000000000000000000000000000", "11": "efa43e0beff151b0f251c4abf48152382b1452b4414dbd737b4127de05ca31f7"}, "127.0.0.1") + status, _ := execRequestAuth(t, s.GetRouter(), http.MethodPost, "/api/l1-builder/v1/register_credentials/rbuilder", sc, nil, measurement.AttestationType, expectedHeader, "127.0.0.1") require.Equal(t, http.StatusOK, status) }) t.Run("CheckCredentials", func(t *testing.T) { resp := make([]ports.BuilderWithServiceCreds, 0) - sc, bts := execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/builders", nil, &resp, measurement.AttestationType, map[string]string{"8": "0000000000000000000000000000000000000000000000000000000000000000", "11": "efa43e0beff151b0f251c4abf48152382b1452b4414dbd737b4127de05ca31f7"}, "127.0.0.1") + sc, bts := execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/builders", nil, &resp, measurement.AttestationType, expectedHeader, "127.0.0.1") c := string(bts) fmt.Println(c) require.Equal(t, http.StatusOK, sc) @@ -228,6 +250,43 @@ func TestAuthInteractionFlow(t *testing.T) { }) } +func TestPortableImageAuthentication(t *testing.T) { + if os.Getenv("RUN_DB_TESTS") != "1" { + t.Skip("skipping test; RUN_DB_TESTS is not set to 1") + } + + s, _, _ := createServer(t) + createBuilder(t, s, "portable-builder", "127.0.0.1", domain.ProductionNetwork) + hashes := &domain.DcapImageHashes{ + UKIAuthenticode: strings.Repeat("1", 96), + KernelAuthenticode: strings.Repeat("2", 96), + CmdlineHash: strings.Repeat("3", 96), + InitrdHash: strings.Repeat("4", 96), + GPTDiskGUIDHash: strings.Repeat("5", 96), + } + measurement := ports.Measurement{ + Name: "portable-image", + AttestationType: "dcap-tdx", + DcapImageHashes: hashes, + } + createMeasurement(t, s, measurement) + + header := map[string]any{ + "type": "image", + "measurements": hashes, + } + var response []ports.BuilderWithServiceCreds + status, _ := execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/builders", nil, &response, measurement.AttestationType, header, "127.0.0.1") + require.Equal(t, http.StatusOK, status) + require.Len(t, response, 1) + + badHashes := *hashes + badHashes.InitrdHash = strings.Repeat("6", 96) + header["measurements"] = &badHashes + status, _ = execRequestAuth(t, s.GetRouter(), http.MethodGet, "/api/l1-builder/v1/builders", nil, &response, measurement.AttestationType, header, "127.0.0.1") + require.Equal(t, http.StatusForbidden, status) +} + // createBuilder emulates admin flow to provision a builder func createBuilder(t *testing.T, s *Server, builderName, ip, network string) { dnsName := builderName + ".builder.net" @@ -349,7 +408,7 @@ func execRequestNoAuth(t *testing.T, router http.Handler, method, url string, re return execRequestAuth(t, router, method, url, request, response, "", nil, "") } -func execRequestAuth(t *testing.T, router http.Handler, method, url string, request, response any, attestationType string, measurement map[string]string, ip string) (statusCode int, responsePayload []byte) { +func execRequestAuth(t *testing.T, router http.Handler, method, url string, request, response any, attestationType string, measurement any, ip string) (statusCode int, responsePayload []byte) { t.Helper() rBody, err := json.Marshal(request) require.NoError(t, err) @@ -358,7 +417,7 @@ func execRequestAuth(t *testing.T, router http.Handler, method, url string, requ if attestationType != "" { tr.Header.Set(ports.AttestationTypeHeader, attestationType) } - if len(measurement) > 0 { + if measurement != nil { measB, err := json.Marshal(measurement) require.NoError(t, err) tr.Header.Set(ports.MeasurementHeader, string(measB)) diff --git a/ports/admin_handler.go b/ports/admin_handler.go index 8644b25..f1a1c45 100644 --- a/ports/admin_handler.go +++ b/ports/admin_handler.go @@ -94,7 +94,12 @@ func (s *AdminHandler) AddMeasurement(w http.ResponseWriter, r *http.Request) { s.BadRequest(w, r, "failed to unmarshal request body", err) return } - err = s.builderService.AddMeasurement(r.Context(), toDomainMeasurement(measurement), false) + domainMeasurement, err := toDomainMeasurement(measurement) + if err != nil { + s.BadRequest(w, r, "invalid measurement", err) + return + } + err = s.builderService.AddMeasurement(r.Context(), domainMeasurement, false) if err != nil { s.log.Error("failed to add measurement", "error", err) w.WriteHeader(http.StatusInternalServerError) diff --git a/ports/http_handler.go b/ports/http_handler.go index 19ad7a9..9adce3c 100644 --- a/ports/http_handler.go +++ b/ports/http_handler.go @@ -18,7 +18,7 @@ import ( type BuilderHubService interface { GetAllowedMeasurements(ctx context.Context) ([]domain.Measurement, error) GetActiveBuilders(ctx context.Context, network string) ([]domain.BuilderWithServices, error) - VerifyIPAndMeasurements(ctx context.Context, ip net.IP, measurement map[string]string, attestationType string) (*domain.Builder, string, error) + VerifyIPAndMeasurements(ctx context.Context, ip net.IP, measurement domain.SuppliedMeasurements, attestationType string) (*domain.Builder, string, error) GetConfigWithSecrets(ctx context.Context, builderName string) ([]byte, error) RegisterCredentialsForBuilder(ctx context.Context, builderName, service, tlsCert string, ecdsaPubKey []byte, measurementName, attestationType, region string) error LogEvent(ctx context.Context, eventName, builderName, name string) error @@ -34,20 +34,108 @@ func NewBuilderHubHandler(builderHubService BuilderHubService, log *httplog.Logg type AuthData struct { AttestationType string - MeasurementData map[string]string + MeasurementData domain.SuppliedMeasurements IP net.IP } +type expectedMeasurementsHeader struct { + Type domain.ExpectedMeasurementType `json:"type"` + Measurements json.RawMessage `json:"measurements"` +} + +func parseMeasurementHeader(value, attestationType string) (domain.SuppliedMeasurements, error) { + var object map[string]json.RawMessage + if err := json.Unmarshal([]byte(value), &object); err != nil { + return domain.SuppliedMeasurements{}, err + } + if object == nil { + return domain.SuppliedMeasurements{}, errors.New("measurement header must be a JSON object") + } + + if _, typed := object["type"]; !typed { + legacy := make(map[string]string) + if err := json.Unmarshal([]byte(value), &legacy); err != nil { + return domain.SuppliedMeasurements{}, err + } + return domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementLegacy, + Registers: nil, + DcapImageHashes: nil, + Legacy: legacy, + }, nil + } + + var header expectedMeasurementsHeader + if err := json.Unmarshal([]byte(value), &header); err != nil { + return domain.SuppliedMeasurements{}, err + } + if !measurementTypeMatchesAttestationType(header.Type, attestationType) { + return domain.SuppliedMeasurements{}, fmt.Errorf("measurement type %q does not match attestation type %q", header.Type, attestationType) + } + + supplied := domain.SuppliedMeasurements{Type: header.Type} + switch header.Type { + case domain.ExpectedMeasurementDCAP, domain.ExpectedMeasurementAzure: + if len(header.Measurements) == 0 || string(header.Measurements) == "null" { + return domain.SuppliedMeasurements{}, errors.New("measurements are missing") + } + if err := json.Unmarshal(header.Measurements, &supplied.Registers); err != nil { + return domain.SuppliedMeasurements{}, err + } + if supplied.Registers == nil { + return domain.SuppliedMeasurements{}, errors.New("measurements must be an object") + } + for key, values := range supplied.Registers { + if len(values) == 0 { + return domain.SuppliedMeasurements{}, fmt.Errorf("measurement %q has no expected values", key) + } + } + case domain.ExpectedMeasurementImage: + if len(header.Measurements) == 0 || string(header.Measurements) == "null" { + return domain.SuppliedMeasurements{}, errors.New("image measurements are missing") + } + supplied.DcapImageHashes = new(domain.DcapImageHashes) + if err := json.Unmarshal(header.Measurements, supplied.DcapImageHashes); err != nil { + return domain.SuppliedMeasurements{}, err + } + if err := supplied.DcapImageHashes.Validate(); err != nil { + return domain.SuppliedMeasurements{}, err + } + case domain.ExpectedMeasurementNoAttestation: + if len(header.Measurements) > 0 && string(header.Measurements) != "null" { + return domain.SuppliedMeasurements{}, errors.New("no_attestation must not contain measurements") + } + default: + return domain.SuppliedMeasurements{}, fmt.Errorf("unknown measurement type %q", header.Type) + } + + return supplied, nil +} + +func measurementTypeMatchesAttestationType(measurementType domain.ExpectedMeasurementType, attestationType string) bool { + switch measurementType { + case domain.ExpectedMeasurementDCAP: + return attestationType == "dcap-tdx" || attestationType == "gcp-tdx" + case domain.ExpectedMeasurementAzure: + return attestationType == "azure-tdx" + case domain.ExpectedMeasurementImage: + return attestationType == "dcap-tdx" || attestationType == "gcp-tdx" + case domain.ExpectedMeasurementNoAttestation: + return attestationType == "none" + default: + return false + } +} + func (bhs *BuilderHubHandler) getAuthData(r *http.Request) (*AuthData, error) { attestationType := r.Header.Get(AttestationTypeHeader) if attestationType == "" { return nil, fmt.Errorf("attestation type is empty %w", ErrInvalidAuthData) } measurementHeader := r.Header.Get(MeasurementHeader) - measurementData := make(map[string]string) - err := json.Unmarshal([]byte(measurementHeader), &measurementData) + measurementData, err := parseMeasurementHeader(measurementHeader, attestationType) if err != nil { - return nil, fmt.Errorf("failed to unmarshal measurement header %w", ErrInvalidAuthData) + return nil, fmt.Errorf("failed to parse measurement header: %v %w", err, ErrInvalidAuthData) } ipHeaders := r.Header.Values(ForwardedHeader) if len(ipHeaders) == 0 { @@ -207,6 +295,11 @@ func (bhs *BuilderHubHandler) GetConfigSecrets(w http.ResponseWriter, r *http.Re w.WriteHeader(http.StatusForbidden) return } + if err != nil { + bhs.log.Error("failed to verify ip and measurements", "error", err) + w.WriteHeader(http.StatusInternalServerError) + return + } bts, err := bhs.builderHubService.GetConfigWithSecrets(r.Context(), builder.Name) if err != nil { bhs.log.Error("failed to get config with secrets", "error", err) diff --git a/ports/types.go b/ports/types.go index fff3ac6..3daf789 100644 --- a/ports/types.go +++ b/ports/types.go @@ -121,7 +121,8 @@ func fromDomainBuilderWithServices(builder domain.BuilderWithServices) BuilderWi type Measurement struct { Name string `json:"measurement_id"` AttestationType string `json:"attestation_type"` - Measurements map[string]domain.SingleMeasurement `json:"measurements"` + Measurements map[string]domain.SingleMeasurement `json:"measurements,omitempty"` + DcapImageHashes *domain.DcapImageHashes `json:"dcap_image_hashes,omitempty"` } func fromDomainMeasurement(measurement domain.Measurement) Measurement { @@ -129,13 +130,35 @@ func fromDomainMeasurement(measurement domain.Measurement) Measurement { Name: measurement.Name, AttestationType: measurement.AttestationType, Measurements: measurement.Measurement, + DcapImageHashes: measurement.DcapImageHashes, } return m } -func toDomainMeasurement(measurement Measurement) domain.Measurement { - m := domain.NewMeasurement(measurement.Name, measurement.AttestationType, measurement.Measurements) - return *m +func toDomainMeasurement(measurement Measurement) (domain.Measurement, error) { + if measurement.Measurements != nil && measurement.DcapImageHashes != nil { + return domain.Measurement{}, errors.New("measurements and dcap_image_hashes are mutually exclusive") + } + if measurement.DcapImageHashes != nil { + if measurement.AttestationType != "dcap-tdx" && measurement.AttestationType != "gcp-tdx" { + return domain.Measurement{}, errors.New("dcap_image_hashes require dcap-tdx or gcp-tdx attestation type") + } + if err := measurement.DcapImageHashes.Validate(); err != nil { + return domain.Measurement{}, err + } + } + + measurements := measurement.Measurements + if measurements == nil && measurement.DcapImageHashes == nil { + measurements = make(map[string]domain.SingleMeasurement) + } + + return domain.Measurement{ + Name: measurement.Name, + AttestationType: measurement.AttestationType, + Measurement: measurements, + DcapImageHashes: measurement.DcapImageHashes, + }, nil } type Builder struct { diff --git a/ports/types_test.go b/ports/types_test.go index 048b867..a575aab 100644 --- a/ports/types_test.go +++ b/ports/types_test.go @@ -2,7 +2,11 @@ package ports import ( "encoding/json" + "strings" "testing" + + "github.com/flashbots/builder-hub/domain" + "github.com/stretchr/testify/require" ) func TestServiceCreds(t *testing.T) { @@ -67,3 +71,109 @@ func TestUnmarshalBuilders(t *testing.T) { t.Error("Failed to unmarshal TLS cert") } } + +func TestParseMeasurementHeader(t *testing.T) { + t.Run("legacy", func(t *testing.T) { + measurement, err := parseMeasurementHeader(`{"4":"aaaa"}`, "custom-attestation") + require.NoError(t, err) + require.Equal(t, domain.ExpectedMeasurementLegacy, measurement.Type) + require.Equal(t, map[string]string{"4": "aaaa"}, measurement.Legacy) + }) + + t.Run("DCAP", func(t *testing.T) { + measurement, err := parseMeasurementHeader(`{"type":"dcap","measurements":{"0":["aaaa"],"3":["bbbb","cccc"]}}`, "dcap-tdx") + require.NoError(t, err) + require.Equal(t, domain.ExpectedMeasurementDCAP, measurement.Type) + require.Equal(t, []string{"bbbb", "cccc"}, measurement.Registers["3"]) + }) + + t.Run("Azure", func(t *testing.T) { + measurement, err := parseMeasurementHeader(`{"type":"azure","measurements":{}}`, "azure-tdx") + require.NoError(t, err) + require.NotNil(t, measurement.Registers) + require.Empty(t, measurement.Registers) + }) + + for _, attestationType := range []string{"dcap-tdx", "gcp-tdx"} { + t.Run("image "+attestationType, func(t *testing.T) { + hashes := validImageHashes() + header, err := json.Marshal(map[string]any{ + "type": "image", + "measurements": hashes, + }) + require.NoError(t, err) + + measurement, err := parseMeasurementHeader(string(header), attestationType) + require.NoError(t, err) + require.Equal(t, &hashes, measurement.DcapImageHashes) + }) + } + + t.Run("no attestation", func(t *testing.T) { + measurement, err := parseMeasurementHeader(`{"type":"no_attestation"}`, "none") + require.NoError(t, err) + require.Equal(t, domain.ExpectedMeasurementNoAttestation, measurement.Type) + }) + + t.Run("typed header does not fall back", func(t *testing.T) { + _, err := parseMeasurementHeader(`{"type":"future","4":"aaaa"}`, "azure-tdx") + require.Error(t, err) + }) + + t.Run("rejects null", func(t *testing.T) { + _, err := parseMeasurementHeader(`null`, "dcap-tdx") + require.Error(t, err) + }) + + t.Run("rejects mismatched attestation type", func(t *testing.T) { + _, err := parseMeasurementHeader(`{"type":"image","measurements":{}}`, "azure-tdx") + require.Error(t, err) + }) + + t.Run("rejects empty register alternatives", func(t *testing.T) { + _, err := parseMeasurementHeader(`{"type":"dcap","measurements":{"0":[]}}`, "dcap-tdx") + require.Error(t, err) + }) + + t.Run("rejects measurements for no attestation", func(t *testing.T) { + _, err := parseMeasurementHeader(`{"type":"no_attestation","measurements":{}}`, "none") + require.Error(t, err) + }) +} + +func TestToDomainMeasurementPortablePolicy(t *testing.T) { + hashes := validImageHashes() + measurement, err := toDomainMeasurement(Measurement{ + Name: "portable", + AttestationType: "dcap-tdx", + DcapImageHashes: &hashes, + }) + require.NoError(t, err) + require.Nil(t, measurement.Measurement) + require.Equal(t, &hashes, measurement.DcapImageHashes) + + _, err = toDomainMeasurement(Measurement{ + Name: "invalid", + AttestationType: "azure-tdx", + DcapImageHashes: &hashes, + }) + require.Error(t, err) + + _, err = toDomainMeasurement(Measurement{ + Name: "both", + AttestationType: "gcp-tdx", + Measurements: map[string]domain.SingleMeasurement{}, + DcapImageHashes: &hashes, + }) + require.Error(t, err) +} + +func validImageHashes() domain.DcapImageHashes { + return domain.DcapImageHashes{ + UKIAuthenticode: strings.Repeat("1", 96), + KernelAuthenticode: strings.Repeat("2", 96), + CmdlineHash: strings.Repeat("3", 96), + InitrdHash: strings.Repeat("4", 96), + GPTDiskGUIDHash: strings.Repeat("5", 96), + } +} diff --git a/schema/005_dcap_image_hashes.sql b/schema/005_dcap_image_hashes.sql new file mode 100644 index 0000000..03f552d --- /dev/null +++ b/schema/005_dcap_image_hashes.sql @@ -0,0 +1,6 @@ +ALTER TABLE measurements_whitelist + ALTER COLUMN measurement DROP NOT NULL, + ADD COLUMN dcap_image_hashes JSONB, + ADD CONSTRAINT measurement_policy_shape CHECK ( + measurement IS NULL OR dcap_image_hashes IS NULL + ); diff --git a/scripts/ci/e2e-test.hurl b/scripts/ci/e2e-test.hurl index 3d3b26b..edb7a15 100644 --- a/scripts/ci/e2e-test.hurl +++ b/scripts/ci/e2e-test.hurl @@ -5,7 +5,7 @@ POST http://localhost:8081/api/admin/v1/measurements { "measurement_id": "test1", - "attestation_type": "test", + "attestation_type": "dcap-tdx", "measurements": {} } HTTP 200 @@ -22,7 +22,7 @@ GET http://localhost:8082/api/l1-builder/v1/measurements HTTP 200 [Asserts] jsonpath "$.[0].measurement_id" == "test1" -jsonpath "$.[0].attestation_type" == "test" +jsonpath "$.[0].attestation_type" == "dcap-tdx" # # BUILDER SETUP diff --git a/testdata/get-measurements.json b/testdata/get-measurements.json index 6ab2361..1b38654 100644 --- a/testdata/get-measurements.json +++ b/testdata/get-measurements.json @@ -52,5 +52,16 @@ "expected": "c9f429296634072d1063a03fb287bed0b2d177b0a22222222222222222222222" } } + }, + { + "measurement_id": "portable-image-v1", + "attestation_type": "dcap-tdx", + "dcap_image_hashes": { + "uki_authenticode": "111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111", + "kernel_authenticode": "222222222222222222222222222222222222222222222222222222222222222222222222222222222222222222222222", + "cmdline_hash": "333333333333333333333333333333333333333333333333333333333333333333333333333333333333333333333333", + "initrd_hash": "444444444444444444444444444444444444444444444444444444444444444444444444444444444444444444444444", + "gpt_disk_guid_hash": "555555555555555555555555555555555555555555555555555555555555555555555555555555555555555555555555" + } } ] From 14a6fc1875d9c14004ea68f2e2f4cca379e4ccfc Mon Sep 17 00:00:00 2001 From: peg Date: Thu, 20 Aug 2026 10:55:35 +0200 Subject: [PATCH 2/9] Add gaurd for legacy measurement values --- application/service.go | 9 ++++++++- application/service_test.go | 16 ++++++++++++++++ 2 files changed, 24 insertions(+), 1 deletion(-) diff --git a/application/service.go b/application/service.go index c3384bc..e0d3d91 100644 --- a/application/service.go +++ b/application/service.go @@ -250,6 +250,13 @@ func equalImageHashes(left, right *domain.DcapImageHashes) bool { // validates that all fields from measurementTemplate are the same in measurement. // For each field, the measurement value must match at least one of the expected values (OR semantics). func checkMeasurement(measurement map[string]string, measurementTemplate domain.Measurement) bool { + // Legacy headers contain observed register values and cannot represent an + // image-hash policy. Without this guard, the loop below could succeed when + // it shouldn't. + if measurementTemplate.DcapImageHashes != nil { + return false + } + for k, v := range measurementTemplate.Measurement { val, ok := measurement[k] if !ok { @@ -262,7 +269,7 @@ func checkMeasurement(measurement map[string]string, measurementTemplate domain. return true } -// matchesAnyExpected returns true if the value matches any of the expected values. +// matchesAnyExpected returns true if the value matches any of the expected values func matchesAnyExpected(value string, expected []string) bool { for _, exp := range expected { if value == exp { diff --git a/application/service_test.go b/application/service_test.go index 912a45e..27049d3 100644 --- a/application/service_test.go +++ b/application/service_test.go @@ -160,6 +160,22 @@ func TestValidateMeasurement(t *testing.T) { _, err := validateMeasurement(legacyMeasurements(measurement), templates) require.ErrorIs(t, err, domain.ErrNotFound) }) + + t.Run("legacy header does not match image policy", func(t *testing.T) { + hashes := testImageHashes("a") + imageTemplates := []domain.Measurement{ + { + Name: "portable-image", + AttestationType: "dcap-tdx", + DcapImageHashes: &hashes, + }, + } + + _, err := validateMeasurement(legacyMeasurements(map[string]string{ + "0": strings.Repeat("a", 96), + }), imageTemplates) + require.ErrorIs(t, err, domain.ErrNotFound) + }) } func TestCheckExpectedMeasurementRegisters(t *testing.T) { From bc769c6ac426d6a0cfe0a15a9ba92d6c11c835e8 Mon Sep 17 00:00:00 2001 From: peg Date: Fri, 21 Aug 2026 08:13:18 +0200 Subject: [PATCH 3/9] Small improvements to README --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 8cd0d04..238a4e2 100644 --- a/README.md +++ b/README.md @@ -218,7 +218,7 @@ For example, a DCAP policy is represented as compact JSON with an array of accep {"type":"dcap","measurements":{"0":[""],"3":[""]}} ``` -Portable image policies use `type: "image"`; `no_attestation` has no `measurements` field. Legacy headers containing a plain register-to-string object remain accepted during migration. +Portable image policies use `"type": "image"`. If no attestation was provided, it will have`"type": "no_attestation"` and no `measurements` field. Legacy headers containing a plain register-to-string object are accepted for backwards compatibility. --- From d3c94b8d4274c92ccc2933b7b12be04d55f68fcc Mon Sep 17 00:00:00 2001 From: peg Date: Mon, 24 Aug 2026 09:41:56 +0200 Subject: [PATCH 4/9] Following copilot review, address issue with handing attestation type: none policies --- application/service.go | 2 +- application/service_test.go | 9 +++++++++ ports/types.go | 3 +++ ports/types_test.go | 18 ++++++++++++++++++ 4 files changed, 31 insertions(+), 1 deletion(-) diff --git a/application/service.go b/application/service.go index e0d3d91..4676068 100644 --- a/application/service.go +++ b/application/service.go @@ -115,7 +115,7 @@ func checkExpectedMeasurement(measurement domain.SuppliedMeasurements, template case domain.ExpectedMeasurementImage: return equalImageHashes(measurement.DcapImageHashes, template.DcapImageHashes) case domain.ExpectedMeasurementNoAttestation: - return true + return len(template.Measurement) == 0 && template.DcapImageHashes == nil default: return false } diff --git a/application/service_test.go b/application/service_test.go index 27049d3..ab697cc 100644 --- a/application/service_test.go +++ b/application/service_test.go @@ -275,6 +275,15 @@ func TestCheckExpectedMeasurementAllowAnyAndNoAttestation(t *testing.T) { require.True(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ Type: domain.ExpectedMeasurementNoAttestation, }, domain.Measurement{AttestationType: "none"})) + + require.False(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementNoAttestation, + }, domain.Measurement{ + AttestationType: "none", + Measurement: map[string]domain.SingleMeasurement{ + "0": {Expected: "aaaa"}, + }, + })) } func legacyMeasurements(values map[string]string) domain.SuppliedMeasurements { diff --git a/ports/types.go b/ports/types.go index 3daf789..0b3721b 100644 --- a/ports/types.go +++ b/ports/types.go @@ -139,6 +139,9 @@ func toDomainMeasurement(measurement Measurement) (domain.Measurement, error) { if measurement.Measurements != nil && measurement.DcapImageHashes != nil { return domain.Measurement{}, errors.New("measurements and dcap_image_hashes are mutually exclusive") } + if measurement.AttestationType == "none" && len(measurement.Measurements) > 0 { + return domain.Measurement{}, errors.New("none attestation type must not contain measurements") + } if measurement.DcapImageHashes != nil { if measurement.AttestationType != "dcap-tdx" && measurement.AttestationType != "gcp-tdx" { return domain.Measurement{}, errors.New("dcap_image_hashes require dcap-tdx or gcp-tdx attestation type") diff --git a/ports/types_test.go b/ports/types_test.go index a575aab..f36d893 100644 --- a/ports/types_test.go +++ b/ports/types_test.go @@ -168,6 +168,24 @@ func TestToDomainMeasurementPortablePolicy(t *testing.T) { require.Error(t, err) } +func TestToDomainMeasurementNoAttestationPolicy(t *testing.T) { + measurement, err := toDomainMeasurement(Measurement{ + Name: "no-attestation", + AttestationType: "none", + }) + require.NoError(t, err) + require.Empty(t, measurement.Measurement) + + _, err = toDomainMeasurement(Measurement{ + Name: "invalid-no-attestation", + AttestationType: "none", + Measurements: map[string]domain.SingleMeasurement{ + "0": {Expected: "aaaa"}, + }, + }) + require.EqualError(t, err, "none attestation type must not contain measurements") +} + func validImageHashes() domain.DcapImageHashes { return domain.DcapImageHashes{ UKIAuthenticode: strings.Repeat("1", 96), From 5ea939348c521b5a59c44e1d17500f75c80abb75 Mon Sep 17 00:00:00 2001 From: peg Date: Mon, 24 Aug 2026 09:48:11 +0200 Subject: [PATCH 5/9] Improve database constraints following Copilot review --- adapters/database/service_test.go | 24 ++++++++++++++++++++++++ application/service.go | 3 +++ application/service_test.go | 14 +++++++++++++- schema/005_dcap_image_hashes.sql | 4 +++- 4 files changed, 43 insertions(+), 2 deletions(-) diff --git a/adapters/database/service_test.go b/adapters/database/service_test.go index 99df124..6209161 100644 --- a/adapters/database/service_test.go +++ b/adapters/database/service_test.go @@ -77,6 +77,30 @@ func TestPortableMeasurement(t *testing.T) { require.Len(t, measurements, 1) require.Nil(t, measurements[0].Measurement) require.Equal(t, hashes, measurements[0].DcapImageHashes) + + t.Run("rejects missing policy representation", func(t *testing.T) { + _, err := dbService.DB.Exec(` + INSERT INTO measurements_whitelist (name, attestation_type, measurement, dcap_image_hashes) + VALUES ('missing-policy', 'dcap-tdx', NULL, NULL) + `) + require.Error(t, err) + }) + + t.Run("rejects JSON null policy representation", func(t *testing.T) { + _, err := dbService.DB.Exec(` + INSERT INTO measurements_whitelist (name, attestation_type, measurement, dcap_image_hashes) + VALUES ('json-null-policy', 'dcap-tdx', 'null'::jsonb, NULL) + `) + require.Error(t, err) + }) + + t.Run("rejects non-object policy representation", func(t *testing.T) { + _, err := dbService.DB.Exec(` + INSERT INTO measurements_whitelist (name, attestation_type, measurement, dcap_image_hashes) + VALUES ('array-policy', 'dcap-tdx', '[]'::jsonb, NULL) + `) + require.Error(t, err) + }) } func TestAdminFlow(t *testing.T) { diff --git a/application/service.go b/application/service.go index 4676068..3fb572e 100644 --- a/application/service.go +++ b/application/service.go @@ -128,6 +128,9 @@ func expectedMeasurementType(template domain.Measurement) (domain.ExpectedMeasur } return "", false } + if template.Measurement == nil { + return "", false + } switch template.AttestationType { case "dcap-tdx", "gcp-tdx": diff --git a/application/service_test.go b/application/service_test.go index ab697cc..a924e03 100644 --- a/application/service_test.go +++ b/application/service_test.go @@ -274,7 +274,10 @@ func TestCheckExpectedMeasurementAllowAnyAndNoAttestation(t *testing.T) { require.True(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ Type: domain.ExpectedMeasurementNoAttestation, - }, domain.Measurement{AttestationType: "none"})) + }, domain.Measurement{ + AttestationType: "none", + Measurement: map[string]domain.SingleMeasurement{}, + })) require.False(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ Type: domain.ExpectedMeasurementNoAttestation, @@ -284,6 +287,15 @@ func TestCheckExpectedMeasurementAllowAnyAndNoAttestation(t *testing.T) { "0": {Expected: "aaaa"}, }, })) + + require.False(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementNoAttestation, + }, domain.Measurement{AttestationType: "none"})) + + require.False(t, checkExpectedMeasurement(domain.SuppliedMeasurements{ + Type: domain.ExpectedMeasurementDCAP, + Registers: map[string][]string{}, + }, domain.Measurement{AttestationType: "dcap-tdx"})) } func legacyMeasurements(values map[string]string) domain.SuppliedMeasurements { diff --git a/schema/005_dcap_image_hashes.sql b/schema/005_dcap_image_hashes.sql index 03f552d..9eb35c1 100644 --- a/schema/005_dcap_image_hashes.sql +++ b/schema/005_dcap_image_hashes.sql @@ -2,5 +2,7 @@ ALTER TABLE measurements_whitelist ALTER COLUMN measurement DROP NOT NULL, ADD COLUMN dcap_image_hashes JSONB, ADD CONSTRAINT measurement_policy_shape CHECK ( - measurement IS NULL OR dcap_image_hashes IS NULL + (measurement IS NULL) <> (dcap_image_hashes IS NULL) + AND (measurement IS NULL OR jsonb_typeof(measurement) = 'object') + AND (dcap_image_hashes IS NULL OR jsonb_typeof(dcap_image_hashes) = 'object') ); From 274909e6873c5164746f92add38db6765aad98f3 Mon Sep 17 00:00:00 2001 From: peg Date: Wed, 26 Aug 2026 09:44:03 +0200 Subject: [PATCH 6/9] Make it harder to accidentally run db tests on production db --- .github/workflows/checks.yml | 5 +-- Makefile | 11 +++++- adapters/database/service_test.go | 17 +++------ httpserver/e2e_test.go | 3 +- internal/testutil/database.go | 57 ++++++++++++++++++++++++++++++ internal/testutil/database_test.go | 54 ++++++++++++++++++++++++++++ 6 files changed, 130 insertions(+), 17 deletions(-) create mode 100644 internal/testutil/database.go create mode 100644 internal/testutil/database_test.go diff --git a/.github/workflows/checks.yml b/.github/workflows/checks.yml index c6b0705..e61df35 100644 --- a/.github/workflows/checks.yml +++ b/.github/workflows/checks.yml @@ -18,6 +18,7 @@ jobs: env: POSTGRES_USER: postgres POSTGRES_PASSWORD: postgres + POSTGRES_DB: builder_hub_test options: >- --health-cmd pg_isready --health-interval 10s @@ -36,7 +37,7 @@ jobs: uses: actions/checkout@v4 - name: Run migrations - run: for file in schema/*.sql; do psql "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable" -f $file; done + run: for file in schema/*.sql; do psql "postgres://postgres:postgres@localhost:5432/builder_hub_test?sslmode=disable" -f $file; done - name: Run unit tests run: make test-with-db @@ -92,4 +93,4 @@ jobs: run: | curl --location --remote-name https://github.com/Orange-OpenSource/hurl/releases/download/6.1.1/hurl_6.1.1_amd64.deb sudo dpkg -i hurl_6.1.1_amd64.deb - ./scripts/ci/integration-test.sh \ No newline at end of file + ./scripts/ci/integration-test.sh diff --git a/Makefile b/Makefile index 29902f6..f29344f 100644 --- a/Makefile +++ b/Makefile @@ -47,6 +47,15 @@ dev-postgres-up: ## Start the PostgreSQL database for development dev-postgres-down: ## Stop the PostgreSQL database for development docker rm -f postgres-test +.PHONY: test-postgres-up +test-postgres-up: ## Start the dedicated disposable PostgreSQL test database + docker run -d --name builder-hub-postgres-test -p 5432:5432 -e POSTGRES_USER=postgres -e POSTGRES_PASSWORD=postgres -e POSTGRES_DB=builder_hub_test postgres + for file in schema/*.sql; do psql "postgres://postgres:postgres@localhost:5432/builder_hub_test?sslmode=disable" -f $file; done + +.PHONY: test-postgres-down +test-postgres-down: ## Stop the dedicated disposable PostgreSQL test database + docker rm -f builder-hub-postgres-test + .PHONY: dev-docker-compose-up dev-docker-compose-up: ## Start Docker compose docker compose -f docker/docker-compose.yaml build @@ -75,7 +84,7 @@ test: ## Run tests .PHONY: test-with-db test-with-db: ## Run tests including live database tests - RUN_DB_TESTS=1 go test -race ./... + RUN_DB_TESTS=1 TEST_POSTGRES_DSN="postgres://postgres:postgres@localhost:5432/builder_hub_test?sslmode=disable" go test -race ./... .PHONY: lint lint: ## Run linters diff --git a/adapters/database/service_test.go b/adapters/database/service_test.go index 6209161..ac2e2a5 100644 --- a/adapters/database/service_test.go +++ b/adapters/database/service_test.go @@ -4,19 +4,16 @@ import ( "context" "encoding/json" "net" - "os" "strings" "testing" "github.com/flashbots/builder-hub/domain" + "github.com/flashbots/builder-hub/internal/testutil" "github.com/stretchr/testify/require" ) func TestGetBuilder(t *testing.T) { - if os.Getenv("RUN_DB_TESTS") != "1" { - t.Skip("skipping test; RUN_DB_TESTS is not set to 1") - } - serv, err := NewDatabaseService("postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable") + serv, err := NewDatabaseService(testutil.PostgresDSN(t)) if err != nil { t.Errorf("NewDatabaseService() = %v; want nil", err) } @@ -50,10 +47,7 @@ func TestGetBuilder(t *testing.T) { } func TestPortableMeasurement(t *testing.T) { - if os.Getenv("RUN_DB_TESTS") != "1" { - t.Skip("skipping test; RUN_DB_TESTS is not set to 1") - } - dbService, err := NewDatabaseService("postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable") + dbService, err := NewDatabaseService(testutil.PostgresDSN(t)) require.NoError(t, err) _, err = dbService.DB.Exec("TRUNCATE TABLE public.measurements_whitelist CASCADE") require.NoError(t, err) @@ -104,10 +98,7 @@ func TestPortableMeasurement(t *testing.T) { } func TestAdminFlow(t *testing.T) { - if os.Getenv("RUN_DB_TESTS") != "1" { - t.Skip("skipping test; RUN_DB_TESTS is not set to 1") - } - dbService, err := NewDatabaseService("postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable") + dbService, err := NewDatabaseService(testutil.PostgresDSN(t)) if err != nil { t.Errorf("NewDatabaseService() = %v; want nil", err) } diff --git a/httpserver/e2e_test.go b/httpserver/e2e_test.go index bb3df31..9318eaf 100644 --- a/httpserver/e2e_test.go +++ b/httpserver/e2e_test.go @@ -16,6 +16,7 @@ import ( "github.com/flashbots/builder-hub/adapters/database" "github.com/flashbots/builder-hub/application" "github.com/flashbots/builder-hub/domain" + "github.com/flashbots/builder-hub/internal/testutil" "github.com/flashbots/builder-hub/ports" "github.com/stretchr/testify/require" ) @@ -371,7 +372,7 @@ func createMeasurement(t *testing.T, s *Server, measurement ports.Measurement) { func createDbService(t *testing.T) *database.Service { t.Helper() - dbService, err := database.NewDatabaseService("postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable") + dbService, err := database.NewDatabaseService(testutil.PostgresDSN(t)) if err != nil { t.Errorf("NewDatabaseService() = %v; want nil", err) } diff --git a/internal/testutil/database.go b/internal/testutil/database.go new file mode 100644 index 0000000..261f711 --- /dev/null +++ b/internal/testutil/database.go @@ -0,0 +1,57 @@ +package testutil + +import ( + "fmt" + "net" + "net/url" + "os" + "strings" + "testing" +) + +const TestDatabaseName = "builder_hub_test" + +// PostgresDSN returns the explicitly configured test database DSN after +// verifying that it points to a dedicated database on the local machine. +// Database tests truncate tables, so they must fail closed if the target is +// not unmistakably disposable. +func PostgresDSN(t testing.TB) string { + t.Helper() + + if os.Getenv("RUN_DB_TESTS") != "1" { + t.Skip("skipping database test; RUN_DB_TESTS is not set to 1") + } + + dsn := strings.TrimSpace(os.Getenv("TEST_POSTGRES_DSN")) + if dsn == "" { + t.Fatal("refusing to run destructive database test: TEST_POSTGRES_DSN is not set") + } + if err := validatePostgresDSN(dsn); err != nil { + t.Fatalf("refusing to run destructive database test: %v", err) + } + + return dsn +} + +func validatePostgresDSN(dsn string) error { + parsed, err := url.Parse(dsn) + if err != nil { + return fmt.Errorf("invalid TEST_POSTGRES_DSN: %w", err) + } + if parsed.Scheme != "postgres" && parsed.Scheme != "postgresql" { + return fmt.Errorf("unsupported TEST_POSTGRES_DSN scheme %q", parsed.Scheme) + } + + host := parsed.Hostname() + ip := net.ParseIP(host) + if !strings.EqualFold(host, "localhost") && (ip == nil || !ip.IsLoopback()) { + return fmt.Errorf("non-loopback host %q", host) + } + + databaseName := strings.TrimPrefix(parsed.Path, "/") + if databaseName != TestDatabaseName { + return fmt.Errorf("database is %q; expected %q", databaseName, TestDatabaseName) + } + + return nil +} diff --git a/internal/testutil/database_test.go b/internal/testutil/database_test.go new file mode 100644 index 0000000..e704619 --- /dev/null +++ b/internal/testutil/database_test.go @@ -0,0 +1,54 @@ +package testutil + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestValidatePostgresDSN(t *testing.T) { + testCases := []struct { + name string + dsn string + wantErr bool + }{ + { + name: "localhost test database", + dsn: "postgres://postgres:postgres@localhost:5432/builder_hub_test?sslmode=disable", + }, + { + name: "IPv4 loopback test database", + dsn: "postgres://postgres:postgres@127.0.0.1:5432/builder_hub_test?sslmode=disable", + }, + { + name: "IPv6 loopback test database", + dsn: "postgres://postgres:postgres@[::1]:5432/builder_hub_test?sslmode=disable", + }, + { + name: "production host", + dsn: "postgres://postgres:postgres@database.example.com:5432/builder_hub_test?sslmode=disable", + wantErr: true, + }, + { + name: "wrong database", + dsn: "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable", + wantErr: true, + }, + { + name: "wrong scheme", + dsn: "mysql://localhost/builder_hub_test", + wantErr: true, + }, + } + + for _, testCase := range testCases { + t.Run(testCase.name, func(t *testing.T) { + err := validatePostgresDSN(testCase.dsn) + if testCase.wantErr { + require.Error(t, err) + return + } + require.NoError(t, err) + }) + } +} From 9f1fc32f1b290dc35019b1b95333a1c1c55aec9f Mon Sep 17 00:00:00 2001 From: peg Date: Wed, 26 Aug 2026 09:46:12 +0200 Subject: [PATCH 7/9] Minor edit to readme --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 238a4e2..6d3a0ad 100644 --- a/README.md +++ b/README.md @@ -218,7 +218,7 @@ For example, a DCAP policy is represented as compact JSON with an array of accep {"type":"dcap","measurements":{"0":[""],"3":[""]}} ``` -Portable image policies use `"type": "image"`. If no attestation was provided, it will have`"type": "no_attestation"` and no `measurements` field. Legacy headers containing a plain register-to-string object are accepted for backwards compatibility. +Portable image policies use `"type": "image"`. If no attestation was provided, it will have `"type": "no_attestation"` and no `measurements` field. Legacy headers containing a plain register-to-string object are accepted for backwards compatibility. --- From dc5fad2ed1581348db8375678180f35f67513e06 Mon Sep 17 00:00:00 2001 From: peg Date: Wed, 26 Aug 2026 10:16:42 +0200 Subject: [PATCH 8/9] Improve guard to avoid running db tests on production db --- adapters/database/service_test.go | 11 ++++----- httpserver/e2e_test.go | 5 ++-- internal/testutil/database.go | 32 +++++++++++++++++++++++++ internal/testutil/database_test.go | 38 ++++++++++++++++++++++++++++++ 4 files changed, 77 insertions(+), 9 deletions(-) diff --git a/adapters/database/service_test.go b/adapters/database/service_test.go index ac2e2a5..61a287f 100644 --- a/adapters/database/service_test.go +++ b/adapters/database/service_test.go @@ -14,9 +14,8 @@ import ( func TestGetBuilder(t *testing.T) { serv, err := NewDatabaseService(testutil.PostgresDSN(t)) - if err != nil { - t.Errorf("NewDatabaseService() = %v; want nil", err) - } + require.NoError(t, err) + testutil.RequireTestDatabase(t, serv.DB) _, err = serv.DB.Exec("TRUNCATE TABLE public.builders CASCADE") require.NoError(t, err) _, err = serv.DB.Exec("TRUNCATE TABLE public.measurements_whitelist CASCADE") @@ -49,6 +48,7 @@ func TestGetBuilder(t *testing.T) { func TestPortableMeasurement(t *testing.T) { dbService, err := NewDatabaseService(testutil.PostgresDSN(t)) require.NoError(t, err) + testutil.RequireTestDatabase(t, dbService.DB) _, err = dbService.DB.Exec("TRUNCATE TABLE public.measurements_whitelist CASCADE") require.NoError(t, err) @@ -99,9 +99,8 @@ func TestPortableMeasurement(t *testing.T) { func TestAdminFlow(t *testing.T) { dbService, err := NewDatabaseService(testutil.PostgresDSN(t)) - if err != nil { - t.Errorf("NewDatabaseService() = %v; want nil", err) - } + require.NoError(t, err) + testutil.RequireTestDatabase(t, dbService.DB) _, err = dbService.DB.Exec("TRUNCATE TABLE public.builders CASCADE") require.NoError(t, err) _, err = dbService.DB.Exec("TRUNCATE TABLE public.measurements_whitelist CASCADE") diff --git a/httpserver/e2e_test.go b/httpserver/e2e_test.go index 9318eaf..2192849 100644 --- a/httpserver/e2e_test.go +++ b/httpserver/e2e_test.go @@ -373,9 +373,8 @@ func createMeasurement(t *testing.T, s *Server, measurement ports.Measurement) { func createDbService(t *testing.T) *database.Service { t.Helper() dbService, err := database.NewDatabaseService(testutil.PostgresDSN(t)) - if err != nil { - t.Errorf("NewDatabaseService() = %v; want nil", err) - } + require.NoError(t, err) + testutil.RequireTestDatabase(t, dbService.DB) _, err = dbService.DB.Exec("TRUNCATE TABLE public.builders CASCADE") require.NoError(t, err) _, err = dbService.DB.Exec("TRUNCATE TABLE public.measurements_whitelist CASCADE") diff --git a/internal/testutil/database.go b/internal/testutil/database.go index 261f711..4c5aeb9 100644 --- a/internal/testutil/database.go +++ b/internal/testutil/database.go @@ -1,3 +1,4 @@ +// Package testutil provides safety helpers for integration tests. package testutil import ( @@ -11,6 +12,10 @@ import ( const TestDatabaseName = "builder_hub_test" +type databaseQueryer interface { + Get(dest any, query string, args ...any) error +} + // PostgresDSN returns the explicitly configured test database DSN after // verifying that it points to a dedicated database on the local machine. // Database tests truncate tables, so they must fail closed if the target is @@ -52,6 +57,33 @@ func validatePostgresDSN(dsn string) error { if databaseName != TestDatabaseName { return fmt.Errorf("database is %q; expected %q", databaseName, TestDatabaseName) } + for parameter := range parsed.Query() { + if parameter != "sslmode" { + return fmt.Errorf("unsupported TEST_POSTGRES_DSN query parameter %q", parameter) + } + } + + return nil +} + +// RequireTestDatabase verifies the connected database identity before a test +// performs destructive operations. +func RequireTestDatabase(t testing.TB, db databaseQueryer) { + t.Helper() + + if err := validateConnectedDatabase(db); err != nil { + t.Fatalf("refusing to run destructive database test: %v", err) + } +} + +func validateConnectedDatabase(db databaseQueryer) error { + var databaseName string + if err := db.Get(&databaseName, "SELECT current_database()"); err != nil { + return fmt.Errorf("failed to verify connected database: %w", err) + } + if databaseName != TestDatabaseName { + return fmt.Errorf("connected database is %q; expected %q", databaseName, TestDatabaseName) + } return nil } diff --git a/internal/testutil/database_test.go b/internal/testutil/database_test.go index e704619..59a75f1 100644 --- a/internal/testutil/database_test.go +++ b/internal/testutil/database_test.go @@ -1,6 +1,7 @@ package testutil import ( + "errors" "testing" "github.com/stretchr/testify/require" @@ -24,6 +25,16 @@ func TestValidatePostgresDSN(t *testing.T) { name: "IPv6 loopback test database", dsn: "postgres://postgres:postgres@[::1]:5432/builder_hub_test?sslmode=disable", }, + { + name: "target overridden in query", + dsn: "postgres://postgres:postgres@localhost:5432/builder_hub_test?host=production.example.com&dbname=production", + wantErr: true, + }, + { + name: "service configuration in query", + dsn: "postgres://postgres:postgres@localhost:5432/builder_hub_test?service=production", + wantErr: true, + }, { name: "production host", dsn: "postgres://postgres:postgres@database.example.com:5432/builder_hub_test?sslmode=disable", @@ -52,3 +63,30 @@ func TestValidatePostgresDSN(t *testing.T) { }) } } + +type stubDatabase struct { + databaseName string + err error +} + +func (s stubDatabase) Get(dest any, _ string, _ ...any) error { + if s.err != nil { + return s.err + } + *(dest.(*string)) = s.databaseName + return nil +} + +func TestValidateConnectedDatabase(t *testing.T) { + t.Run("accepts test database", func(t *testing.T) { + require.NoError(t, validateConnectedDatabase(stubDatabase{databaseName: TestDatabaseName})) + }) + + t.Run("rejects another database", func(t *testing.T) { + require.Error(t, validateConnectedDatabase(stubDatabase{databaseName: "production"})) + }) + + t.Run("rejects verification failure", func(t *testing.T) { + require.Error(t, validateConnectedDatabase(stubDatabase{err: errors.New("query failed")})) + }) +} From 4f8f0eeda09f0f855cc98c414c1338aa44fd5f4e Mon Sep 17 00:00:00 2001 From: peg Date: Wed, 26 Aug 2026 10:40:18 +0200 Subject: [PATCH 9/9] Use a response-specific measurement type --- ports/http_handler.go | 2 +- ports/types.go | 19 ++++++++++++++++--- ports/types_test.go | 35 +++++++++++++++++++++++++++++++++++ 3 files changed, 52 insertions(+), 4 deletions(-) diff --git a/ports/http_handler.go b/ports/http_handler.go index 9adce3c..17c6916 100644 --- a/ports/http_handler.go +++ b/ports/http_handler.go @@ -172,7 +172,7 @@ func (bhs *BuilderHubHandler) GetAllowedMeasurements(w http.ResponseWriter, r *h w.WriteHeader(http.StatusInternalServerError) return } - pMeasurements := make([]Measurement, 0, len(measurements)) + pMeasurements := make([]measurementResponse, 0, len(measurements)) for _, m := range measurements { pMeasurements = append(pMeasurements, fromDomainMeasurement(m)) } diff --git a/ports/types.go b/ports/types.go index 0b3721b..1283268 100644 --- a/ports/types.go +++ b/ports/types.go @@ -125,13 +125,26 @@ type Measurement struct { DcapImageHashes *domain.DcapImageHashes `json:"dcap_image_hashes,omitempty"` } -func fromDomainMeasurement(measurement domain.Measurement) Measurement { - m := Measurement{ +type measurementResponse struct { + Name string `json:"measurement_id"` + AttestationType string `json:"attestation_type"` + Measurements *map[string]domain.SingleMeasurement `json:"measurements,omitempty"` + DcapImageHashes *domain.DcapImageHashes `json:"dcap_image_hashes,omitempty"` +} + +func fromDomainMeasurement(measurement domain.Measurement) measurementResponse { + m := measurementResponse{ Name: measurement.Name, AttestationType: measurement.AttestationType, - Measurements: measurement.Measurement, DcapImageHashes: measurement.DcapImageHashes, } + if measurement.DcapImageHashes == nil { + measurements := measurement.Measurement + if measurements == nil { + measurements = make(map[string]domain.SingleMeasurement) + } + m.Measurements = &measurements + } return m } diff --git a/ports/types_test.go b/ports/types_test.go index f36d893..db7cdbe 100644 --- a/ports/types_test.go +++ b/ports/types_test.go @@ -186,6 +186,41 @@ func TestToDomainMeasurementNoAttestationPolicy(t *testing.T) { require.EqualError(t, err, "none attestation type must not contain measurements") } +func TestMeasurementResponseJSON(t *testing.T) { + t.Run("preserves empty register policy", func(t *testing.T) { + response := fromDomainMeasurement(domain.Measurement{ + Name: "allow-any", + AttestationType: "dcap-tdx", + Measurement: map[string]domain.SingleMeasurement{}, + }) + + encoded, err := json.Marshal(response) + require.NoError(t, err) + require.JSONEq(t, `{ + "measurement_id": "allow-any", + "attestation_type": "dcap-tdx", + "measurements": {} + }`, string(encoded)) + }) + + t.Run("omits register policy for image policy", func(t *testing.T) { + hashes := validImageHashes() + response := fromDomainMeasurement(domain.Measurement{ + Name: "portable", + AttestationType: "gcp-tdx", + DcapImageHashes: &hashes, + }) + + encoded, err := json.Marshal(response) + require.NoError(t, err) + + var object map[string]json.RawMessage + require.NoError(t, json.Unmarshal(encoded, &object)) + require.NotContains(t, object, "measurements") + require.Contains(t, object, "dcap_image_hashes") + }) +} + func validImageHashes() domain.DcapImageHashes { return domain.DcapImageHashes{ UKIAuthenticode: strings.Repeat("1", 96),