Fix LDAP review feedback
This commit is contained in:
Firehawk 2026-07-13 09:36:26 +09:30 committed by GitHub
commit ad92603bb1
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 29 additions and 21 deletions

View File

@ -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)
}
}

View File

@ -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)
}

View File

@ -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
}

View File

@ -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"))
})
})