Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 17 additions & 54 deletions service/entityresolution/multi-strategy/registration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,12 @@ import (
"github.com/opentdf/platform/protocol/go/entityresolution"
"github.com/opentdf/platform/service/entityresolution/multi-strategy/types"
"github.com/opentdf/platform/service/logger"
"github.com/stretchr/testify/require"
"google.golang.org/protobuf/types/known/anypb"
"google.golang.org/protobuf/types/known/structpb"
)

func TestResolveEntities_ClaimsProviderUsesInlineClaimsContext(t *testing.T) {
t.Helper()

erService, err := NewERS(t.Context(), types.MultiStrategyConfig{
Providers: map[string]types.ProviderConfig{
"jwt": {
Expand Down Expand Up @@ -48,22 +47,16 @@ func TestResolveEntities_ClaimsProviderUsesInlineClaimsContext(t *testing.T) {
},
},
}, logger.CreateTestLogger())
if err != nil {
t.Fatalf("NewERS() error = %v", err)
}
require.NoError(t, err)

claimsStruct, err := structpb.NewStruct(map[string]interface{}{
"sub": "diana",
"email": "diana@example.com",
})
if err != nil {
t.Fatalf("structpb.NewStruct() error = %v", err)
}
require.NoError(t, err)

claimsAny, err := anypb.New(claimsStruct)
if err != nil {
t.Fatalf("anypb.New() error = %v", err)
}
require.NoError(t, err)

resp, err := erService.ResolveEntities(t.Context(), connect.NewRequest(&entityresolution.ResolveEntitiesRequest{
Entities: []*authorization.Entity{
Expand All @@ -73,37 +66,20 @@ func TestResolveEntities_ClaimsProviderUsesInlineClaimsContext(t *testing.T) {
},
},
}))
if err != nil {
t.Fatalf("ResolveEntities() error = %v", err)
}

if got := len(resp.Msg.GetEntityRepresentations()); got != 1 {
t.Fatalf("expected 1 entity representation, got %d", got)
}
require.NoError(t, err)
require.Len(t, resp.Msg.GetEntityRepresentations(), 1)

props := resp.Msg.GetEntityRepresentations()[0].GetAdditionalProps()
if len(props) != 1 {
t.Fatalf("expected 1 additional props entry, got %d", len(props))
}
require.Len(t, props, 1)

result := props[0].AsMap()
if got := result["subject"]; got != "diana" {
t.Fatalf("expected subject diana, got %v", got)
}
if got := result["email_address"]; got != "diana@example.com" {
t.Fatalf("expected email_address diana@example.com, got %v", got)
}
if got := result["metadata_source"]; got != "jwt_claims" {
t.Fatalf("expected metadata_source jwt_claims, got %v", got)
}
if _, hasError := result["error"]; hasError {
t.Fatalf("expected successful resolution, got error payload: %v", result["error"])
}
require.Equal(t, "diana", result["subject"])
require.Equal(t, "diana@example.com", result["email_address"])
require.Equal(t, "jwt_claims", result["metadata_source"])
require.NotContains(t, result, "error")
}

func TestResolveEntities_UserNameEntityDoesNotSeedClaimsContext(t *testing.T) {
t.Helper()

erService, err := NewERS(t.Context(), types.MultiStrategyConfig{
Providers: map[string]types.ProviderConfig{
"jwt": {
Expand All @@ -124,34 +100,21 @@ func TestResolveEntities_UserNameEntityDoesNotSeedClaimsContext(t *testing.T) {
},
},
}, logger.CreateTestLogger())
if err != nil {
t.Fatalf("NewERS() error = %v", err)
}
require.NoError(t, err)

resp, err := erService.ResolveEntities(t.Context(), connect.NewRequest(&entityresolution.ResolveEntitiesRequest{
Entities: []*authorization.Entity{{
Id: "alice-user-name",
EntityType: &authorization.Entity_UserName{UserName: "alice"},
}},
}))
if err != nil {
t.Fatalf("ResolveEntities() error = %v", err)
}

if got := len(resp.Msg.GetEntityRepresentations()); got != 1 {
t.Fatalf("expected 1 entity representation, got %d", got)
}
require.NoError(t, err)
require.Len(t, resp.Msg.GetEntityRepresentations(), 1)

props := resp.Msg.GetEntityRepresentations()[0].GetAdditionalProps()
if len(props) != 1 {
t.Fatalf("expected 1 additional props entry, got %d", len(props))
}
require.Len(t, props, 1)

result := props[0].AsMap()
if _, hasError := result["error"]; !hasError {
t.Fatalf("expected claims provider to fail without middleware claims for user_name entity, got %v", result)
}
if got := result["entity_id"]; got != "alice-user-name" {
t.Fatalf("expected entity_id alice-user-name, got %v", got)
}
require.Contains(t, result, "error")
require.Equal(t, "alice-user-name", result["entity_id"])
}
118 changes: 28 additions & 90 deletions service/entityresolution/multi-strategy/v2/registration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -112,21 +112,15 @@ func TestERSV2_ResolveEntities_PopulatesRepresentations(t *testing.T) {
}

ers, err := NewERSV2(t.Context(), config, logger.CreateTestLogger())
if err != nil {
t.Fatalf("Failed to create ERSV2: %v", err)
}
require.NoError(t, err)

claimsStruct, err := structpb.NewStruct(map[string]interface{}{
"sub": "alice",
"email": "alice@example.com",
})
if err != nil {
t.Fatalf("Failed to build claims struct: %v", err)
}
require.NoError(t, err)
claimsAny, err := anypb.New(claimsStruct)
if err != nil {
t.Fatalf("Failed to wrap claims in anypb.Any: %v", err)
}
require.NoError(t, err)

req := connect.NewRequest(&ersV2.ResolveEntitiesRequest{
Entities: []*entity.Entity{
Expand All @@ -138,52 +132,32 @@ func TestERSV2_ResolveEntities_PopulatesRepresentations(t *testing.T) {
})

resp, err := ers.ResolveEntities(t.Context(), req)
if err != nil {
t.Fatalf("ResolveEntities returned error: %v", err)
}
require.NoError(t, err)

reps := resp.Msg.GetEntityRepresentations()
if len(reps) != 1 {
t.Fatalf("EntityRepresentations length = %d, want 1 (empty response means the handler silently dropped the entity via structpb.NewStruct failure)", len(reps))
}
if got := reps[0].GetOriginalId(); got != "entity-1" {
t.Errorf("OriginalId = %q, want %q", got, "entity-1")
}
require.Len(t, reps, 1, "empty response means the handler silently dropped the entity via structpb.NewStruct failure")
require.Equal(t, "entity-1", reps[0].GetOriginalId())

props := reps[0].GetAdditionalProps()
if len(props) != 1 {
t.Fatalf("AdditionalProps length = %d, want 1", len(props))
}
require.Len(t, props, 1)
fields := props[0].GetFields()

// The resolved claim should be present.
if got := fields["username"].GetStringValue(); got != "alice" {
t.Errorf("username in AdditionalProps = %q, want %q", got, "alice")
}
require.Equal(t, "alice", fields["username"].GetStringValue())

// metadata_attempted_strategies MUST serialize to a ListValue. If the
// source-level fix regresses and the field is stored as []string again,
// structpb.NewStruct will drop the whole entity and this assertion (and
// the length assertion above) will fail.
metaAttempted, ok := fields["metadata_attempted_strategies"]
if !ok {
t.Fatalf("metadata_attempted_strategies missing from AdditionalProps; the handler likely dropped the entity")
}
require.True(t, ok, "metadata_attempted_strategies missing from AdditionalProps; the handler likely dropped the entity")
list := metaAttempted.GetListValue()
if list == nil {
t.Fatalf("metadata_attempted_strategies must be a ListValue, got kind %T", metaAttempted.GetKind())
}
if got, want := len(list.GetValues()), 1; got != want {
t.Errorf("metadata_attempted_strategies length = %d, want %d", got, want)
}
if got := list.GetValues()[0].GetStringValue(); got != "jwt_strategy" {
t.Errorf("metadata_attempted_strategies[0] = %q, want %q", got, "jwt_strategy")
}
require.NotNil(t, list, "metadata_attempted_strategies has kind %T", metaAttempted.GetKind())
require.Len(t, list.GetValues(), 1)
require.Equal(t, "jwt_strategy", list.GetValues()[0].GetStringValue())
}

func TestResolveEntities_ClaimsProviderUsesInlineClaimsContext(t *testing.T) {
t.Helper()

erService, err := NewERSV2(t.Context(), types.MultiStrategyConfig{
Providers: map[string]types.ProviderConfig{
"jwt": {
Expand Down Expand Up @@ -217,22 +191,16 @@ func TestResolveEntities_ClaimsProviderUsesInlineClaimsContext(t *testing.T) {
},
},
}, logger.CreateTestLogger())
if err != nil {
t.Fatalf("NewERSV2() error = %v", err)
}
require.NoError(t, err)

claimsStruct, err := structpb.NewStruct(map[string]interface{}{
"sub": "diana",
"email": "diana@example.com",
})
if err != nil {
t.Fatalf("structpb.NewStruct() error = %v", err)
}
require.NoError(t, err)

claimsAny, err := anypb.New(claimsStruct)
if err != nil {
t.Fatalf("anypb.New() error = %v", err)
}
require.NoError(t, err)

resp, err := erService.ResolveEntities(t.Context(), connect.NewRequest(&ersV2.ResolveEntitiesRequest{
Entities: []*entity.Entity{
Expand All @@ -243,37 +211,20 @@ func TestResolveEntities_ClaimsProviderUsesInlineClaimsContext(t *testing.T) {
},
},
}))
if err != nil {
t.Fatalf("ResolveEntities() error = %v", err)
}

if got := len(resp.Msg.GetEntityRepresentations()); got != 1 {
t.Fatalf("expected 1 entity representation, got %d", got)
}
require.NoError(t, err)
require.Len(t, resp.Msg.GetEntityRepresentations(), 1)

props := resp.Msg.GetEntityRepresentations()[0].GetAdditionalProps()
if len(props) != 1 {
t.Fatalf("expected 1 additional props entry, got %d", len(props))
}
require.Len(t, props, 1)

result := props[0].AsMap()
if got := result["subject"]; got != "diana" {
t.Fatalf("expected subject diana, got %v", got)
}
if got := result["email_address"]; got != "diana@example.com" {
t.Fatalf("expected email_address diana@example.com, got %v", got)
}
if got := result["metadata_source"]; got != "jwt_claims" {
t.Fatalf("expected metadata_source jwt_claims, got %v", got)
}
if _, hasError := result["error"]; hasError {
t.Fatalf("expected successful resolution, got error payload: %v", result["error"])
}
require.Equal(t, "diana", result["subject"])
require.Equal(t, "diana@example.com", result["email_address"])
require.Equal(t, "jwt_claims", result["metadata_source"])
require.NotContains(t, result, "error")
}

func TestResolveEntities_UserNameEntityDoesNotSeedClaimsContext(t *testing.T) {
t.Helper()

erService, err := NewERSV2(t.Context(), types.MultiStrategyConfig{
Providers: map[string]types.ProviderConfig{
"jwt": {
Expand All @@ -294,9 +245,7 @@ func TestResolveEntities_UserNameEntityDoesNotSeedClaimsContext(t *testing.T) {
},
},
}, logger.CreateTestLogger())
if err != nil {
t.Fatalf("NewERSV2() error = %v", err)
}
require.NoError(t, err)

resp, err := erService.ResolveEntities(t.Context(), connect.NewRequest(&ersV2.ResolveEntitiesRequest{
Entities: []*entity.Entity{{
Expand All @@ -305,26 +254,15 @@ func TestResolveEntities_UserNameEntityDoesNotSeedClaimsContext(t *testing.T) {
Category: entity.Entity_CATEGORY_SUBJECT,
}},
}))
if err != nil {
t.Fatalf("ResolveEntities() error = %v", err)
}

if got := len(resp.Msg.GetEntityRepresentations()); got != 1 {
t.Fatalf("expected 1 entity representation, got %d", got)
}
require.NoError(t, err)
require.Len(t, resp.Msg.GetEntityRepresentations(), 1)

props := resp.Msg.GetEntityRepresentations()[0].GetAdditionalProps()
if len(props) != 1 {
t.Fatalf("expected 1 additional props entry, got %d", len(props))
}
require.Len(t, props, 1)

result := props[0].AsMap()
if _, hasError := result["error"]; !hasError {
t.Fatalf("expected claims provider to fail without middleware claims for user_name entity, got %v", result)
}
if got := result["entity_id"]; got != "alice-user-name" {
t.Fatalf("expected entity_id alice-user-name, got %v", got)
}
require.Contains(t, result, "error")
require.Equal(t, "alice-user-name", result["entity_id"])
}

func TestCreateEntityFromResultV2ExcludesResolutionMetadataFromPolicyClaims(t *testing.T) {
Expand Down
22 changes: 16 additions & 6 deletions service/pkg/protohelper/structpb.go
Original file line number Diff line number Diff line change
@@ -1,16 +1,18 @@
package protohelper

// StructPBCompatibleValue normalizes Go values into shapes accepted by structpb.NewStruct.
// In particular, it recursively converts []string into []interface{} while preserving
// In particular, it recursively converts typed slices into []interface{} while preserving
// nested []interface{} and map[string]interface{} values.
func StructPBCompatibleValue(value interface{}) interface{} {
switch v := value.(type) {
case []string:
result := make([]interface{}, len(v))
for i, item := range v {
result[i] = item
}
return result
return structPBSlice(v)
case []float64:
return structPBSlice(v)
case []bool:
return structPBSlice(v)
case []int:
return structPBSlice(v)
case []interface{}:
result := make([]interface{}, len(v))
for i, item := range v {
Expand All @@ -27,3 +29,11 @@ func StructPBCompatibleValue(value interface{}) interface{} {
return value
}
}

func structPBSlice[T any](values []T) []interface{} {
result := make([]interface{}, len(values))
for i, value := range values {
result[i] = StructPBCompatibleValue(value)
}
return result
}
Loading
Loading