@@ -66,18 +66,6 @@ func ExtractEmailFromAssertion(assertion *saml.Assertion) (string, error) {
|
|||||||
return "", fmt.Errorf("could not extract email from assertion")
|
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 {
|
func MapSAMLRoleToSystemRole(samlRole string) *coredata.MembershipRole {
|
||||||
if samlRole != "" && isValidRole(samlRole) {
|
if samlRole != "" && isValidRole(samlRole) {
|
||||||
role := coredata.MembershipRole(samlRole)
|
role := coredata.MembershipRole(samlRole)
|
||||||
@@ -104,7 +92,7 @@ func ExtractUserAttributes(
|
|||||||
if assertion.Subject != nil && assertion.Subject.NameID != nil {
|
if assertion.Subject != nil && assertion.Subject.NameID != nil {
|
||||||
email, err = mail.ParseAddr(assertion.Subject.NameID.Value)
|
email, err = mail.ParseAddr(assertion.Subject.NameID.Value)
|
||||||
if err != nil {
|
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()
|
fullname = email.String()
|
||||||
role = ""
|
role = ""
|
||||||
@@ -123,7 +111,7 @@ func ExtractUserAttributes(
|
|||||||
}
|
}
|
||||||
email, err = mail.ParseAddr(emailString)
|
email, err = mail.ParseAddr(emailString)
|
||||||
if err != nil {
|
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)
|
firstname, err := ExtractAttributeValue(assertion, attributeFirstname)
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/crewjam/saml"
|
"github.com/crewjam/saml"
|
||||||
@@ -517,7 +518,7 @@ func (s *SAMLService) HandleSAMLAssertion(
|
|||||||
return nil, ErrCannotExtractUserAttributes{Err: err}
|
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)
|
return nil, fmt.Errorf("email domain mismatch: assertion contains email with domain %s but SAML config is for domain %s", email.Domain(), config.EmailDomain)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -17,7 +17,16 @@ func (a Addr) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *Addr) Domain() 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) {
|
func ParseAddr(s string) (Addr, error) {
|
||||||
@@ -39,6 +48,11 @@ func (a Addr) Value() (driver.Value, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *Addr) Scan(value any) error {
|
func (a *Addr) Scan(value any) error {
|
||||||
|
if value == nil {
|
||||||
|
*a = Nil
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
switch v := value.(type) {
|
switch v := value.(type) {
|
||||||
case string:
|
case string:
|
||||||
parsed, err := ParseAddr(v)
|
parsed, err := ParseAddr(v)
|
||||||
|
|||||||
@@ -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.TrustCenterID, "trust_center_id", validator.Required(), validator.GID(coredata.TrustCenterEntityType))
|
||||||
v.Check(ctcar.Email, "email", validator.Required(), validator.NotEmpty())
|
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))
|
v.Check(ctcar.Name, "name", validator.SafeTextNoNewLine(TitleMaxLength))
|
||||||
|
|
||||||
return v.Error()
|
return v.Error()
|
||||||
|
|||||||
@@ -2022,8 +2022,10 @@ func (r *mutationResolver) UpdatePeople(ctx context.Context, input types.UpdateP
|
|||||||
|
|
||||||
var emailAddresses []mail.Addr
|
var emailAddresses []mail.Addr
|
||||||
for _, emailAddress := range input.AdditionalEmailAddresses {
|
for _, emailAddress := range input.AdditionalEmailAddresses {
|
||||||
|
if emailAddress != nil {
|
||||||
emailAddresses = append(emailAddresses, *emailAddress)
|
emailAddresses = append(emailAddresses, *emailAddress)
|
||||||
}
|
}
|
||||||
|
}
|
||||||
people, err := prb.Peoples.Update(ctx, probo.UpdatePeopleRequest{
|
people, err := prb.Peoples.Update(ctx, probo.UpdatePeopleRequest{
|
||||||
ID: input.ID,
|
ID: input.ID,
|
||||||
FullName: input.FullName,
|
FullName: input.FullName,
|
||||||
|
|||||||
@@ -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")
|
return nil, fmt.Errorf("email and name are not allowed for authenticated users")
|
||||||
}
|
}
|
||||||
|
|
||||||
*email = tokenData.Email
|
email = &tokenData.Email
|
||||||
}
|
}
|
||||||
if email == nil {
|
if email == nil {
|
||||||
return nil, fmt.Errorf("email is required for unauthenticated users")
|
return nil, fmt.Errorf("email is required for unauthenticated users")
|
||||||
|
|||||||
@@ -199,7 +199,7 @@ func BenchmarkPattern(b *testing.B) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func BenchmarkValidate_WithErrors(b *testing.B) {
|
func BenchmarkValidate_WithErrors(b *testing.B) {
|
||||||
email := "invalid-email"
|
email := ""
|
||||||
|
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
|
|||||||
@@ -27,13 +27,12 @@ var (
|
|||||||
strings.Split(strings.TrimSpace(string(disposableEmailsRaw)), "\n"),
|
strings.Split(strings.TrimSpace(string(disposableEmailsRaw)), "\n"),
|
||||||
testEmails...,
|
testEmails...,
|
||||||
)
|
)
|
||||||
|
notOneOfBlacklisted = NotOneOfSlice(blacklistedEmails)
|
||||||
)
|
)
|
||||||
|
|
||||||
func NotBlacklisted() ValidatorFunc {
|
func NotBlacklisted() ValidatorFunc {
|
||||||
notOneOfSlice := NotOneOfSlice(blacklistedEmails)
|
|
||||||
|
|
||||||
return func(value any) *ValidationError {
|
return func(value any) *ValidationError {
|
||||||
err := notOneOfSlice(value)
|
err := notOneOfBlacklisted(value)
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return newValidationError(
|
return newValidationError(
|
||||||
|
|||||||
@@ -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.
|
// 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.
|
// Accepts a slice of any type. Compares by value first, then by string representation.
|
||||||
func NotOneOfSlice[T any](disallowed []T) ValidatorFunc {
|
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 {
|
return func(value any) *ValidationError {
|
||||||
err := oneOfSlice(value)
|
// Handle nil values first
|
||||||
|
if value == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
if err == nil {
|
// Dereference all pointer levels
|
||||||
return newValidationError(
|
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,
|
ErrorCodeInvalidEnum,
|
||||||
"must not be one of the disallowed values",
|
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,
|
||||||
|
fmt.Sprintf("must not be one of: %s", strings.Join(disallowedStrings, ", ")),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user