diff --git a/pkg/auth/saml_mapper.go b/pkg/auth/saml_mapper.go index 612aa19e4..3ef21882f 100644 --- a/pkg/auth/saml_mapper.go +++ b/pkg/auth/saml_mapper.go @@ -66,18 +66,6 @@ func ExtractEmailFromAssertion(assertion *saml.Assertion) (string, error) { return "", fmt.Errorf("could not extract email from assertion") } -func ExtractEmailDomain(email string) (string, error) { - parts := strings.Split(email, "@") - if len(parts) != 2 { - return "", fmt.Errorf("invalid email address: %s", email) - } - domain := strings.ToLower(strings.TrimSpace(parts[1])) - if domain == "" { - return "", fmt.Errorf("empty domain in email address: %s", email) - } - return domain, nil -} - func MapSAMLRoleToSystemRole(samlRole string) *coredata.MembershipRole { if samlRole != "" && isValidRole(samlRole) { role := coredata.MembershipRole(samlRole) @@ -104,7 +92,7 @@ func ExtractUserAttributes( if assertion.Subject != nil && assertion.Subject.NameID != nil { email, err = mail.ParseAddr(assertion.Subject.NameID.Value) if err != nil { - return mail.Nil, "", "", fmt.Errorf("invalid nameID as email address") + return mail.Nil, "", "", fmt.Errorf("invalid nameID as email address: %w", err) } fullname = email.String() role = "" @@ -123,7 +111,7 @@ func ExtractUserAttributes( } email, err = mail.ParseAddr(emailString) if err != nil { - return mail.Nil, "", "", fmt.Errorf("invalid attribute email") + return mail.Nil, "", "", fmt.Errorf("invalid attribute email: %w", err) } firstname, err := ExtractAttributeValue(assertion, attributeFirstname) diff --git a/pkg/auth/saml_service.go b/pkg/auth/saml_service.go index 0c19ccecb..c2456763a 100644 --- a/pkg/auth/saml_service.go +++ b/pkg/auth/saml_service.go @@ -24,6 +24,7 @@ import ( "fmt" "net/http" "net/url" + "strings" "time" "github.com/crewjam/saml" @@ -517,7 +518,7 @@ func (s *SAMLService) HandleSAMLAssertion( return nil, ErrCannotExtractUserAttributes{Err: err} } - if email.Domain() != config.EmailDomain { + if !strings.EqualFold(email.Domain(), config.EmailDomain) { return nil, fmt.Errorf("email domain mismatch: assertion contains email with domain %s but SAML config is for domain %s", email.Domain(), config.EmailDomain) } diff --git a/pkg/mail/addr.go b/pkg/mail/addr.go index eb08815cb..f8923138e 100644 --- a/pkg/mail/addr.go +++ b/pkg/mail/addr.go @@ -17,7 +17,16 @@ func (a Addr) String() string { } func (a *Addr) Domain() string { - return strings.Split(a.String(), "@")[1] + if a == nil || *a == Nil { + return "" + } + + parts := strings.Split(a.String(), "@") + if len(parts) != 2 { + return "" + } + + return parts[1] } func ParseAddr(s string) (Addr, error) { @@ -39,6 +48,11 @@ func (a Addr) Value() (driver.Value, error) { } func (a *Addr) Scan(value any) error { + if value == nil { + *a = Nil + return nil + } + switch v := value.(type) { case string: parsed, err := ParseAddr(v) diff --git a/pkg/probo/trust_center_access_service.go b/pkg/probo/trust_center_access_service.go index dd20bbb5a..0905e0444 100644 --- a/pkg/probo/trust_center_access_service.go +++ b/pkg/probo/trust_center_access_service.go @@ -66,6 +66,7 @@ func (ctcar *CreateTrustCenterAccessRequest) Validate() error { v.Check(ctcar.TrustCenterID, "trust_center_id", validator.Required(), validator.GID(coredata.TrustCenterEntityType)) v.Check(ctcar.Email, "email", validator.Required(), validator.NotEmpty()) + v.Check(ctcar.Email.Domain(), "email", validator.NotBlacklisted()) v.Check(ctcar.Name, "name", validator.SafeTextNoNewLine(TitleMaxLength)) return v.Error() diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index f55f68a3a..dc3eec5db 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -2022,7 +2022,9 @@ func (r *mutationResolver) UpdatePeople(ctx context.Context, input types.UpdateP var emailAddresses []mail.Addr for _, emailAddress := range input.AdditionalEmailAddresses { - emailAddresses = append(emailAddresses, *emailAddress) + if emailAddress != nil { + emailAddresses = append(emailAddresses, *emailAddress) + } } people, err := prb.Peoples.Update(ctx, probo.UpdatePeopleRequest{ ID: input.ID, diff --git a/pkg/server/api/trust/v1/v1_resolver.go b/pkg/server/api/trust/v1/v1_resolver.go index b4b0ed7dd..7f9a2a9cf 100644 --- a/pkg/server/api/trust/v1/v1_resolver.go +++ b/pkg/server/api/trust/v1/v1_resolver.go @@ -135,7 +135,7 @@ func (r *mutationResolver) RequestAllAccesses(ctx context.Context, input types.R return nil, fmt.Errorf("email and name are not allowed for authenticated users") } - *email = tokenData.Email + email = &tokenData.Email } if email == nil { return nil, fmt.Errorf("email is required for unauthenticated users") diff --git a/pkg/validator/validation_bench_test.go b/pkg/validator/validation_bench_test.go index 615f46094..19dd14e09 100644 --- a/pkg/validator/validation_bench_test.go +++ b/pkg/validator/validation_bench_test.go @@ -199,7 +199,7 @@ func BenchmarkPattern(b *testing.B) { } func BenchmarkValidate_WithErrors(b *testing.B) { - email := "invalid-email" + email := "" b.ResetTimer() for i := 0; i < b.N; i++ { diff --git a/pkg/validator/validator_email.go b/pkg/validator/validator_email.go index 4b7fcbde0..a33074a43 100644 --- a/pkg/validator/validator_email.go +++ b/pkg/validator/validator_email.go @@ -27,13 +27,12 @@ var ( strings.Split(strings.TrimSpace(string(disposableEmailsRaw)), "\n"), testEmails..., ) + notOneOfBlacklisted = NotOneOfSlice(blacklistedEmails) ) func NotBlacklisted() ValidatorFunc { - notOneOfSlice := NotOneOfSlice(blacklistedEmails) - return func(value any) *ValidationError { - err := notOneOfSlice(value) + err := notOneOfBlacklisted(value) if err != nil { return newValidationError( diff --git a/pkg/validator/validator_string.go b/pkg/validator/validator_string.go index 2a350f348..30e152f7f 100644 --- a/pkg/validator/validator_string.go +++ b/pkg/validator/validator_string.go @@ -228,15 +228,49 @@ func OneOfSlice[T any](allowed []T) ValidatorFunc { // NotOneOfSlice validates that a value is not one of the values in the slice. // Accepts a slice of any type. Compares by value first, then by string representation. func NotOneOfSlice[T any](disallowed []T) ValidatorFunc { - oneOfSlice := OneOfSlice(disallowed) + // Build disallowed map with string keys for flexible comparison + disallowedMap := make(map[string]bool) + disallowedStrings := make([]string, 0, len(disallowed)) + + for _, v := range disallowed { + str := fmt.Sprint(v) + disallowedMap[str] = true + disallowedStrings = append(disallowedStrings, str) + } return func(value any) *ValidationError { - err := oneOfSlice(value) + // Handle nil values first + if value == nil { + return nil + } - if err == nil { - return newValidationError( + // Dereference all pointer levels + actualValue := value + val := reflect.ValueOf(value) + for val.Kind() == reflect.Ptr { + if val.IsNil() { + return nil + } + val = val.Elem() + actualValue = val.Interface() + } + + // First try exact match with DeepEqual + for _, disallowedVal := range disallowed { + if reflect.DeepEqual(actualValue, disallowedVal) { + newValidationError( + ErrorCodeInvalidEnum, + fmt.Sprintf("must not be one of: %s", strings.Join(disallowedStrings, ", ")), + ) + } + } + + // Then try string comparison (for custom string types) + valueStr := fmt.Sprint(actualValue) + if disallowedMap[valueStr] { + newValidationError( ErrorCodeInvalidEnum, - "must not be one of the disallowed values", + fmt.Sprintf("must not be one of: %s", strings.Join(disallowedStrings, ", ")), ) }