package auth import ( "context" "errors" "strings" "testing" "time" "github.com/go-ldap/ldap/v3" ) // fakeConn ist ein LDAP-Server-Ersatz. bindErr bildet Bind-Ergebnisse pro // DN ab; searches protokolliert alle Filter für Injection-Assertions. type fakeConn struct { bindErr map[string]error entries map[string][]*ldap.Entry // Filter-Substring -> Ergebnis searches []string closed bool dialErr error } func (f *fakeConn) Bind(dn, pw string) error { if err, ok := f.bindErr[dn+"|"+pw]; ok { return err } if err, ok := f.bindErr[dn]; ok { return err } return nil } func (f *fakeConn) Search(req *ldap.SearchRequest) (*ldap.SearchResult, error) { f.searches = append(f.searches, req.Filter) // Echtes AD vergleicht sAMAccountName ohne Rücksicht auf Groß-/Kleinschreibung; // der Fake muss das nachbilden, sonst testet er strenger als die Wirklichkeit. filter := strings.ToLower(req.Filter) for needle, entries := range f.entries { if strings.Contains(filter, strings.ToLower(needle)) { return &ldap.SearchResult{Entries: entries}, nil } } return &ldap.SearchResult{}, nil } func (f *fakeConn) Close() error { f.closed = true; return nil } func entry(dn string, attrs map[string][]string) *ldap.Entry { e := &ldap.Entry{DN: dn} for name, vals := range attrs { e.Attributes = append(e.Attributes, &ldap.EntryAttribute{Name: name, Values: vals}) } return e } const ( userDN = "CN=Max Mueller,OU=Users,DC=firma,DC=local" groupDN = "CN=VPN-Users,OU=Groups,DC=firma,DC=local" ) // stdEntries liefert die Standardantworten: Gruppen-Auflösung, User-Suche, // Mitgliedschaftsprüfung. func stdEntries(inGroup bool) map[string][]*ldap.Entry { m := map[string][]*ldap.Entry{ "objectClass=group": {entry(groupDN, nil)}, "sAMAccountName=mmueller)(userPrincipalName": { entry(userDN, map[string][]string{"sAMAccountName": {"MMueller"}}), }, } if inGroup { m["1.2.840.113556.1.4.1941"] = []*ldap.Entry{entry(userDN, nil)} } return m } // dialSequence liefert eine Dial-Funktion, die die übergebenen Verbindungen // der Reihe nach ausgibt und danach bei der letzten bleibt. func dialSequence(conns []*fakeConn) func(context.Context, string) (conn, error) { i := 0 return func(ctx context.Context, server string) (conn, error) { c := conns[min(i, len(conns)-1)] i++ if c.dialErr != nil { return nil, c.dialErr } return c, nil } } func baseOptions() Options { return Options{ Servers: []string{"dc01.firma.local", "dc02.firma.local"}, Port: 636, TLSMode: "ldaps", BaseDN: "DC=firma,DC=local", BindUser: "svc@firma.local", BindPassword: "svc-pw", VPNGroup: "VPN-Users", Timeout: time.Second, } } func newAD(t *testing.T, conns ...*fakeConn) *AD { t.Helper() opts := baseOptions() opts.Dial = dialSequence(conns) a, err := NewAD(opts) if err != nil { t.Fatalf("NewAD: %v", err) } return a } func TestAuthenticateSuccessCanonicalisesUsername(t *testing.T) { c := &fakeConn{entries: stdEntries(true)} a := newAD(t, c) id, err := a.Authenticate(context.Background(), " MMueller ", "geheim") if err != nil { t.Fatalf("Authenticate: %v", err) } // Kanonisch = sAMAccountName aus dem Verzeichnis, lowercase. if id.Username != "mmueller" { t.Errorf("Username = %q, want mmueller", id.Username) } if !id.HasGroup("VPN-Users") { t.Errorf("Groups = %v, muss die VPN-Gruppe enthalten", id.Groups) } } func TestAuthenticateRejectsNonMember(t *testing.T) { c := &fakeConn{entries: stdEntries(false)} a := newAD(t, c) _, err := a.Authenticate(context.Background(), "mmueller", "geheim") var ae *Error if !errors.As(err, &ae) || ae.Reason != ReasonNotInVPNGroup { t.Fatalf("err = %v, want Reason %q", err, ReasonNotInVPNGroup) } } func TestAuthenticateMapsBindFailure(t *testing.T) { c := &fakeConn{entries: stdEntries(true)} c.bindErr = map[string]error{ userDN + "|falsch": errors.New( "LDAP Result Code 49 \"Invalid Credentials\": AcceptSecurityContext error, data 532, v4563"), } a := newAD(t, c) _, err := a.Authenticate(context.Background(), "mmueller", "falsch") var ae *Error if !errors.As(err, &ae) || ae.Reason != ReasonPasswordExpired { t.Fatalf("err = %v, want Reason %q", err, ReasonPasswordExpired) } } func TestAuthenticateUnknownUserReportsGenericReason(t *testing.T) { c := &fakeConn{entries: map[string][]*ldap.Entry{"objectClass=group": {entry(groupDN, nil)}}} a := newAD(t, c) _, err := a.Authenticate(context.Background(), "gibtsnicht", "egal") var ae *Error if !errors.As(err, &ae) || ae.Reason != ReasonUserNotFound { t.Fatalf("err = %v, want Reason %q", err, ReasonUserNotFound) } if ae.UserVisible() { t.Error("unbekannter Benutzer darf keine spezifische Meldung erzeugen") } } func TestAuthenticateRejectsEmptyPassword(t *testing.T) { c := &fakeConn{entries: stdEntries(true)} a := newAD(t, c) // Leeres Passwort wäre ein anonymer Bind und würde fälschlich gelingen. _, err := a.Authenticate(context.Background(), "mmueller", "") var ae *Error if !errors.As(err, &ae) || ae.Reason != ReasonInvalidCredentials { t.Fatalf("err = %v, want Reason %q", err, ReasonInvalidCredentials) } for _, f := range c.searches { if strings.Contains(f, "mmueller") { t.Error("bei leerem Passwort darf gar keine Suche stattfinden") } } } func TestAuthenticateEscapesFilterInput(t *testing.T) { c := &fakeConn{entries: stdEntries(true)} a := newAD(t, c) _, _ = a.Authenticate(context.Background(), "evil)(objectClass=*", "pw") for _, f := range c.searches { if strings.Contains(f, "evil)(objectClass=*") { t.Fatalf("unescapte Benutzereingabe im Filter: %q", f) } } var found bool for _, f := range c.searches { if strings.Contains(strings.ToLower(f), `\29\28`) { // ")(" escaped found = true } } if !found { t.Errorf("Eingabe wurde nicht escapt; Filter: %v", c.searches) } } func TestFailoverToSecondDC(t *testing.T) { dead := &fakeConn{dialErr: errors.New("connection refused")} alive := &fakeConn{entries: stdEntries(true)} var failedOver []string opts := baseOptions() opts.OnFailover = func(server string, err error) { failedOver = append(failedOver, server) } opts.Dial = dialSequence([]*fakeConn{dead, alive}) a, err := NewAD(opts) if err != nil { t.Fatal(err) } if _, err := a.Authenticate(context.Background(), "mmueller", "geheim"); err != nil { t.Fatalf("Failover muss gelingen: %v", err) } if len(failedOver) == 0 || failedOver[0] != "dc01.firma.local" { t.Errorf("OnFailover = %v, want dc01.firma.local", failedOver) } } func TestAllDCsDownIsBackendUnavailable(t *testing.T) { dead := &fakeConn{dialErr: errors.New("connection refused")} a := newAD(t, dead) _, err := a.Authenticate(context.Background(), "mmueller", "geheim") var ae *Error if !errors.As(err, &ae) || ae.Reason != ReasonBackendUnavailable { t.Fatalf("err = %v, want Reason %q", err, ReasonBackendUnavailable) } } func TestLookupForTestAuth(t *testing.T) { c := &fakeConn{entries: stdEntries(true)} a := newAD(t, c) res, err := a.Lookup(context.Background(), "MMueller") if err != nil { t.Fatalf("Lookup: %v", err) } if res.DN != userDN || res.SAMAccountName != "mmueller" || !res.InVPNGroup { t.Errorf("Lookup = %+v", res) } } func TestConnectionsAreAlwaysClosed(t *testing.T) { c := &fakeConn{entries: stdEntries(true)} a := newAD(t, c) if _, err := a.Authenticate(context.Background(), "mmueller", "geheim"); err != nil { t.Fatal(err) } if !c.closed { t.Error("Verbindung muss geschlossen werden") } } func TestSetFailoverHookIsUsed(t *testing.T) { dead := &fakeConn{dialErr: errors.New("connection refused")} alive := &fakeConn{entries: stdEntries(true)} opts := baseOptions() opts.Dial = dialSequence([]*fakeConn{dead, alive}) a, err := NewAD(opts) if err != nil { t.Fatal(err) } var seen []string a.SetFailoverHook(func(server string, err error) { seen = append(seen, server) }) if _, err := a.Authenticate(context.Background(), "mmueller", "geheim"); err != nil { t.Fatal(err) } if len(seen) == 0 || seen[0] != "dc01.firma.local" { t.Fatalf("Failover-Hook = %v", seen) } } func TestServiceBindFailureIsNotMaskedByFailover(t *testing.T) { // Ein falsches Dienstkonto-Passwort ist ein Konfigurationsfehler. // Es darf nicht als "nächster DC probieren" durchgehen. c := &fakeConn{ entries: stdEntries(true), bindErr: map[string]error{"svc@firma.local": errors.New("LDAP Result Code 49")}, } a := newAD(t, c) _, err := a.Authenticate(context.Background(), "mmueller", "geheim") if err == nil { t.Fatal("fehlerhafter Service-Bind muss zum Fehler führen") } if !strings.Contains(err.Error(), "Service-Bind") { t.Errorf("Fehler muss den Service-Bind benennen: %v", err) } }