Files
ldaps/server_search.go

253 lines
8.4 KiB
Go

package ldaps
import (
"errors"
"fmt"
"log"
"net"
"runtime/debug"
"strings"
ber "github.com/go-asn1-ber/asn1-ber"
"github.com/go-ldap/ldap/v3"
)
func HandleSearchRequest(req *ber.Packet, controls *[]ldap.Control, messageID uint64, boundDN string, server *Server, conn net.Conn) (resultErr error) {
defer func() {
if r := recover(); r != nil {
log.Printf("Recovered from panic in SearchFn: %s\n%s", r, string(debug.Stack()))
resultErr = ldap.NewError(ldap.LDAPResultOperationsError, fmt.Errorf("Search function panic: %s", r))
}
}()
searchReq, err := parseSearchRequest(boundDN, req, controls)
if err != nil {
log.Printf("Error parsing search request: %v", req.Children[1].Value)
return err
}
filterPacket, err := ldap.CompileFilter(searchReq.Filter)
if err != nil {
log.Printf("Error compiling filter: %v", searchReq.Filter)
return ldap.NewError(ldap.LDAPResultFilterError, err)
}
fnNames := []string{}
for k := range server.SearchFns {
fnNames = append(fnNames, k)
}
fn := routeFunc(searchReq.BaseDN, fnNames)
searchResp, err := server.SearchFns[fn].Search(boundDN, searchReq, conn)
if err != nil {
log.Printf("SearchFn Error %s", err.Error())
return ldap.NewError(searchResp.ResultCode, err)
}
if server.EnforceLDAP {
if searchReq.DerefAliases != ldap.NeverDerefAliases { // [-a {never|always|search|find}
// TODO: Server DerefAliases not supported: RFC4511 4.5.1.3
}
if searchReq.TimeLimit > 0 {
// TODO: Server TimeLimit not implemented
}
}
i := 0
for _, entry := range searchResp.Entries {
if server.EnforceLDAP {
// filter
matched, resultCode := ServerApplyFilter(filterPacket, entry)
// Per https://datatracker.ietf.org/doc/html/rfc4511#section-4.5.1.7
// unimplemented search tags/filters should match as "UNDEFINED" which essentially means they match as false.
// They should not return an error
if resultCode != ldap.LDAPResultSuccess && resultCode != ldap.LDAPResultFilterError {
return ldap.NewError(resultCode, errors.New("ServerApplyFilter error"))
}
if !matched {
continue
}
// constrained search scope
switch searchReq.Scope {
case ldap.ScopeWholeSubtree: // The scope is constrained to the entry named by baseObject and to all its subordinates.
case ldap.ScopeBaseObject: // The scope is constrained to the entry named by baseObject.
if !strings.EqualFold(entry.DN, searchReq.BaseDN) {
continue
}
case ldap.ScopeSingleLevel: // The scope is constrained to the immediate subordinates of the entry named by baseObject.
parts := strings.Split(entry.DN, ",")
if len(parts) < 2 && !strings.EqualFold(entry.DN, searchReq.BaseDN) {
continue
}
if dnSuffix := strings.Join(parts[1:], ","); !strings.EqualFold(dnSuffix, searchReq.BaseDN) {
continue
}
}
// filter attributes
entry, err = filterAttributes(entry, searchReq.Attributes)
if err != nil {
return ldap.NewError(ldap.LDAPResultOperationsError, err)
}
// size limit
if searchReq.SizeLimit > 0 && i >= searchReq.SizeLimit {
break
}
i++
}
// respond
responsePacket := encodeSearchResponse(messageID, entry)
if err = sendPacket(conn, responsePacket); err != nil {
log.Printf("Error encoding response: %v", searchReq.Filter)
return ldap.NewError(ldap.LDAPResultOperationsError, err)
}
}
// If we had a paging control, we need to update its cookie if present
for _, reqcontrol := range *controls {
if reqcontrol.GetControlType() == ldap.ControlTypePaging {
for _, respcontrol := range searchResp.Controls {
if respcontrol.GetControlType() == ldap.ControlTypePaging {
reqcontrol.(*ldap.ControlPaging).Cookie = respcontrol.(*ldap.ControlPaging).Cookie
break
}
}
}
}
return nil
}
func parseSearchRequest(boundDN string, req *ber.Packet, controls *[]ldap.Control) (ldap.SearchRequest, error) {
if len(req.Children) != 8 {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultOperationsError, errors.New("Bad search request"))
}
// Parse the request
baseObject, ok := req.Children[0].Value.(string)
if !ok {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultProtocolError, errors.New("Bad search request"))
}
s, ok := req.Children[1].Value.(int64)
if !ok {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultProtocolError, errors.New("Bad search request"))
}
scope := int(s)
d, ok := req.Children[2].Value.(int64)
if !ok {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultProtocolError, errors.New("Bad search request"))
}
derefAliases := int(d)
s, ok = req.Children[3].Value.(int64)
if !ok {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultProtocolError, errors.New("Bad search request"))
}
sizeLimit := int(s)
t, ok := req.Children[4].Value.(int64)
if !ok {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultProtocolError, errors.New("Bad search request"))
}
timeLimit := int(t)
typesOnly := false
if req.Children[5].Value != nil {
typesOnly, ok = req.Children[5].Value.(bool)
if !ok {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultProtocolError, errors.New("Bad search request"))
}
}
filter, err := ldap.DecompileFilter(req.Children[6])
if err != nil {
return ldap.SearchRequest{}, err
}
attributes := []string{}
for _, attr := range req.Children[7].Children {
a, ok := attr.Value.(string)
if !ok {
return ldap.SearchRequest{}, ldap.NewError(ldap.LDAPResultProtocolError, errors.New("Bad search request"))
}
attributes = append(attributes, a)
}
searchReq := ldap.SearchRequest{
BaseDN: baseObject,
Scope: scope,
DerefAliases: derefAliases,
SizeLimit: sizeLimit,
TimeLimit: timeLimit,
TypesOnly: typesOnly,
Filter: filter,
Attributes: attributes,
Controls: *controls,
}
return searchReq, nil
}
func filterAttributes(entry *ldap.Entry, attributes []string) (*ldap.Entry, error) {
// only return requested attributes
newAttributes := []*ldap.EntryAttribute{}
if len(attributes) > 1 || (len(attributes) == 1 && len(attributes[0]) > 0) {
for _, attr := range entry.Attributes {
for _, requested := range attributes {
// You can request the directory server to return operational attributes by adding + (the plus sign) in your ldapsearch command.
// "+supportedControl" is treated as an operational attribute
if strings.HasPrefix(attr.Name, "+") {
if requested == "+" || strings.EqualFold(attr.Name, "+"+requested) {
newAttributes = append(newAttributes, ldap.NewEntryAttribute(attr.Name[1:], attr.Values))
break
}
} else {
if requested == "*" || strings.EqualFold(attr.Name, requested) {
newAttributes = append(newAttributes, attr)
break
}
}
}
}
} else {
// remove operational attributes
for _, attr := range entry.Attributes {
if !strings.HasPrefix(attr.Name, "+") {
newAttributes = append(newAttributes, attr)
}
}
}
entry.Attributes = newAttributes
return entry, nil
}
func encodeSearchResponse(messageID uint64, res *ldap.Entry) *ber.Packet {
responsePacket := ber.Encode(ber.ClassUniversal, ber.TypeConstructed, ber.TagSequence, nil, "LDAP Response")
responsePacket.AppendChild(ber.NewInteger(ber.ClassUniversal, ber.TypePrimitive, ber.TagInteger, messageID, "Message ID"))
searchEntry := ber.Encode(ber.ClassApplication, ber.TypeConstructed, ldap.ApplicationSearchResultEntry, nil, "Search Result Entry")
searchEntry.AppendChild(ber.NewString(ber.ClassUniversal, ber.TypePrimitive, ber.TagOctetString, res.DN, "Object Name"))
attrs := ber.Encode(ber.ClassUniversal, ber.TypeConstructed, ber.TagSequence, nil, "Attributes:")
for _, attribute := range res.Attributes {
attrs.AppendChild(encodeSearchAttribute(attribute.Name, attribute.Values))
}
searchEntry.AppendChild(attrs)
responsePacket.AppendChild(searchEntry)
return responsePacket
}
func encodeSearchAttribute(name string, values []string) *ber.Packet {
packet := ber.Encode(ber.ClassUniversal, ber.TypeConstructed, ber.TagSequence, nil, "Attribute")
packet.AppendChild(ber.NewString(ber.ClassUniversal, ber.TypePrimitive, ber.TagOctetString, name, "Attribute Name"))
valuesPacket := ber.Encode(ber.ClassUniversal, ber.TypeConstructed, ber.TagSet, nil, "Attribute Values")
for _, value := range values {
valuesPacket.AppendChild(ber.NewString(ber.ClassUniversal, ber.TypePrimitive, ber.TagOctetString, value, "Attribute Value"))
}
packet.AppendChild(valuesPacket)
return packet
}