diff --git a/core/ldapauth/ldapauth.go b/core/ldapauth/ldapauth.go index 9af22d4a5..030dff047 100644 --- a/core/ldapauth/ldapauth.go +++ b/core/ldapauth/ldapauth.go @@ -124,9 +124,11 @@ func (s *Store) Load(ctx context.Context) (Config, error) { } for i := range c.Sources { if c.Sources[i].BindPassword != "" { - if p, e := utils.Decrypt(ctx, s.key, c.Sources[i].BindPassword); e == nil { - c.Sources[i].BindPassword = p + p, e := utils.Decrypt(ctx, s.key, c.Sources[i].BindPassword) + if e != nil { + return c, fmt.Errorf("failed to decrypt bind password for source %s: %w", c.Sources[i].Name, e) } + c.Sources[i].BindPassword = p } applyDefaults(&c.Sources[i]) } @@ -258,7 +260,7 @@ func authLDAP(ctx context.Context, ds model.DataStore, src Source, username, pas func lookupAndBind(src Source, username, password string) (DiscoveredUser, error) { l, err := ldap.DialURL(src.URL) if err != nil { - return DiscoveredUser{}, ErrUserNotFound + return DiscoveredUser{}, err } defer l.Close() if src.StartTLS { @@ -290,10 +292,10 @@ func lookupAndBind(src Source, username, password string) (DiscoveredUser, error return DiscoveredUser{}, ErrBadPassword } du := DiscoveredUser{DN: e.DN, UserName: first(e.GetAttributeValue(src.UserNameAttribute), username), Name: e.GetAttributeValue(src.DisplayNameAttribute), Email: e.GetAttributeValue(src.EmailAttribute), Groups: e.GetAttributeValues("memberOf")} - du.Groups = append(du.Groups, groupsForUser(src, e.DN)...) + du.Groups = append(du.Groups, groupsForUser(src, e.DN, du.UserName)...) return du, nil } -func groupsForUser(src Source, userDN string) []string { +func groupsForUser(src Source, userDN, username string) []string { if src.GroupBaseDN == "" { return nil } @@ -303,10 +305,14 @@ func groupsForUser(src Source, userDN string) []string { } defer l.Close() if src.StartTLS { - _ = l.StartTLS(&tls.Config{InsecureSkipVerify: src.InsecureSkipVerify}) //nolint:gosec + if err = l.StartTLS(&tls.Config{InsecureSkipVerify: src.InsecureSkipVerify}); err != nil { //nolint:gosec + return nil + } } if src.BindDN != "" { - _ = l.Bind(src.BindDN, src.BindPassword) + if err = l.Bind(src.BindDN, src.BindPassword); err != nil { + return nil + } } gf := src.GroupFilter if gf == "" { @@ -319,7 +325,8 @@ func groupsForUser(src Source, userDN string) []string { } var out []string for _, g := range res.Entries { - if slices.Contains(g.GetAttributeValues(src.GroupMemberAttribute), userDN) { + members := g.GetAttributeValues(src.GroupMemberAttribute) + if slices.Contains(members, userDN) || slices.Contains(members, username) { out = append(out, g.DN) } } @@ -353,11 +360,8 @@ func loginUserFilter(src Source, username string) string { if src.UserFilter == "" { return fmt.Sprintf("(%s=%s)", attr, escapedUsername) } - if placeholders := strings.Count(src.UserFilter, "%s"); placeholders > 0 { - if placeholders == 1 { - return fmt.Sprintf(src.UserFilter, escapedUsername) - } - return fmt.Sprintf(src.UserFilter, attr, escapedUsername) + if strings.Contains(src.UserFilter, "%s") { + return strings.ReplaceAll(src.UserFilter, "%s", escapedUsername) } return fmt.Sprintf("(&%s(%s=%s))", src.UserFilter, attr, escapedUsername) } @@ -447,7 +451,7 @@ func TestAndCache(ctx context.Context, src Source) (Source, error) { } for i, u := range src.Cache.Users { for g, m := range grps { - if slices.Contains(m, u.DN) { + if slices.Contains(m, u.DN) || slices.Contains(m, u.UserName) { src.Cache.Users[i].Groups = append(src.Cache.Users[i].Groups, g) } } diff --git a/core/ldapauth/ldapauth_test.go b/core/ldapauth/ldapauth_test.go index 76ef52e74..b6f27d294 100644 --- a/core/ldapauth/ldapauth_test.go +++ b/core/ldapauth/ldapauth_test.go @@ -19,10 +19,10 @@ func TestLoginUserFilter(t *testing.T) { } }) - t.Run("preserves explicit placeholder filters", func(t *testing.T) { + t.Run("replaces every username placeholder", func(t *testing.T) { t.Parallel() - got := loginUserFilter(Source{UserNameAttribute: "uid", UserFilter: "(&(objectClass=person)(%s=%s))"}, "directory-user") - want := "(&(objectClass=person)(uid=directory-user))" + got := loginUserFilter(Source{UserNameAttribute: "uid", UserFilter: "(|(uid=%s)(mail=%s))"}, "directory-user") + want := "(|(uid=directory-user)(mail=directory-user))" if got != want { t.Fatalf("loginUserFilter() = %q, want %q", got, want) } diff --git a/persistence/user_repository.go b/persistence/user_repository.go index abf762f95..e996e7c32 100644 --- a/persistence/user_repository.go +++ b/persistence/user_repository.go @@ -182,6 +182,10 @@ func (r *userRepository) preserveExternalAuthSource(u *model.User) error { if existing.AuthSource == "" { return nil } + if u.AuthSource == "" { + u.AuthSource = existing.AuthSource + u.AuthSourceID = existing.AuthSourceID + } if u.ExternalSync { return nil } @@ -201,10 +205,6 @@ func (r *userRepository) preserveExternalAuthSource(u *model.User) error { if len(validation.Errors) > 0 { return validation } - if u.AuthSource == "" { - u.AuthSource = existing.AuthSource - u.AuthSourceID = existing.AuthSourceID - } return nil } diff --git a/persistence/user_repository_test.go b/persistence/user_repository_test.go index e30b92e5c..88d4d9127 100644 --- a/persistence/user_repository_test.go +++ b/persistence/user_repository_test.go @@ -103,12 +103,16 @@ var _ = Describe("UserRepository", func() { actual.Name = "Synced User" actual.Email = "synced@example.com" + actual.AuthSource = "" + actual.AuthSourceID = "" actual.ExternalSync = true Expect(repo.Put(actual)).To(Succeed()) actual, err = repo.FindByUsername("ldap_user") Expect(err).ToNot(HaveOccurred()) Expect(actual.Name).To(Equal("Synced User")) Expect(actual.Email).To(Equal("synced@example.com")) + Expect(actual.AuthSource).To(Equal("ldap")) + Expect(actual.AuthSourceID).To(Equal("ldap01")) }) })