Ajout config ldapAttr + msg erreur si inexistant dnas LDAP

This commit is contained in:
yo
2021-01-12 13:07:01 +01:00
parent d9a3cae44d
commit 91a0c84593
+159 -147
View File
@@ -1,7 +1,7 @@
// mxrouter // mxrouter
// Copyright (c) 2020 yo000 <johan@nosd.in> // Copyright (c) 2020 yo000 <johan@nosd.in>
// //
// // TODO : REconnexion ldap lorsque necessaire
// checkmx : postfix tcp_table server which returns second-level domain name of the MX of a domain // checkmx : postfix tcp_table server which returns second-level domain name of the MX of a domain
// ex : // ex :
@@ -13,179 +13,191 @@
package main package main
import ( import (
"flag" "flag"
"fmt" "fmt"
"log" "log"
"log/syslog" "log/syslog"
"net" "net"
"os" "os"
"strings" "strings"
"sync" "sync"
"github.com/go-ldap/ldap/v3" "github.com/go-ldap/ldap/v3"
"github.com/peterbourgon/ff" "github.com/peterbourgon/ff"
) )
const ( const (
version = "0.9.1" version = "0.9.3"
) )
var ( var (
logstream *syslog.Writer logstream *syslog.Writer
mutex sync.Mutex mutex sync.Mutex
debug *bool debug *bool
listen *string listen *string
ldapURL *string ldapURL *string
ldapBaseDN *string ldapBaseDN *string
ldapUser *string ldapUser *string
ldapPass *string ldapPass *string
defDomain *string ldapAttr *string
defDomain *string
) )
func handleConnection(connClt net.Conn, conLdap *ldap.Conn) { func handleConnection(connClt net.Conn, conLdap *ldap.Conn) {
var s string var s string
var mxdom string var mxdom string
buf := make([]byte, 1024) buf := make([]byte, 1024)
defer connClt.Close() defer connClt.Close()
readlen, err := connClt.Read(buf) readlen, err := connClt.Read(buf)
if err != nil { if err != nil {
fmt.Println("Error reading:", err.Error()) fmt.Println("Error reading:", err.Error())
} }
// 1) Recupere le MX du domaine de l'adresse mail recue // 1) Recupere le MX du domaine de l'adresse mail recue
if strings.HasPrefix(string(buf[:readlen-1]), "get ") && strings.Contains(string(buf[:readlen-1]), "@") { if strings.HasPrefix(string(buf[:readlen-1]), "get ") && strings.Contains(string(buf[:readlen-1]), "@") {
mail := string(buf[4:readlen-1]) mail := string(buf[4 : readlen-1])
domain := strings.Split(mail, "@")[1] domain := strings.Split(mail, "@")[1]
mxs, err := net.LookupMX(domain) mxs, err := net.LookupMX(domain)
if err != nil { if err != nil {
logstream.Err(fmt.Sprintln("Error lookup mx: ", err)) logstream.Err(fmt.Sprintln("Error lookup mx: ", err))
s := strings.Replace(fmt.Sprintf("Error lookup MX: %s", err.Error()), " ", "%20", -1) s := strings.Replace(fmt.Sprintf("Error lookup MX: %s", err.Error()), " ", "%20", -1)
response := fmt.Sprintf("500 %s\n", s) response := fmt.Sprintf("500 %s\n", s)
connClt.Write([]byte(response)) connClt.Write([]byte(response))
return return
} }
// 2) Requete LDAP pour voir si le domaine de second niveau du MX possede un routage particulier // 2) Requete LDAP pour voir si le domaine de second niveau du MX possede un routage particulier
// Protection contre ca : example.org. 72 IN MX 0 . // Protection contre ca : example.org. 72 IN MX 0 .
// Considerons qu'il n'y aura jamais de "one letter TLD" // Considerons qu'il n'y aura jamais de "one letter TLD"
if len(mxs[0].Host) < 4 { if len(mxs[0].Host) < 4 {
logstream.Err(fmt.Sprintln("No usable mx found")) logstream.Err(fmt.Sprintln("No usable mx found"))
s := strings.Replace("No usable MX found", " ", "%20", -1) s := strings.Replace("No usable MX found", " ", "%20", -1)
resp := fmt.Sprintf("500 %s\n", s) resp := fmt.Sprintf("500 %s\n", s)
connClt.Write([]byte(resp)) connClt.Write([]byte(resp))
return return
} }
mxslic := strings.Split(mxs[0].Host, ".") mxslic := strings.Split(mxs[0].Host, ".")
// les MXs terminent par '.', on le retire // les MXs terminent par '.', on le retire
if mxs[0].Host[len(mxs[0].Host)-1] == '.' { if mxs[0].Host[len(mxs[0].Host)-1] == '.' {
mxdom = strings.Join(mxslic[len(mxslic)-3:len(mxslic)-1], ".") mxdom = strings.Join(mxslic[len(mxslic)-3:len(mxslic)-1], ".")
} else { } else {
mxdom = strings.Join(mxslic[len(mxslic)-2:len(mxslic)], ".") mxdom = strings.Join(mxslic[len(mxslic)-2:len(mxslic)], ".")
} }
filter := fmt.Sprintf("(dc=%s)", ldap.EscapeFilter(mxdom)) filter := fmt.Sprintf("(dc=%s)", ldap.EscapeFilter(mxdom))
searchReq := ldap.NewSearchRequest(*ldapBaseDN, ldap.ScopeWholeSubtree, 0, 0, 0, searchReq := ldap.NewSearchRequest(*ldapBaseDN, ldap.ScopeWholeSubtree, 0, 0, 0,
false, filter, []string{"relayName"}, []ldap.Control{}) false, filter, []string{*ldapAttr}, []ldap.Control{})
mutex.Lock() mutex.Lock()
result, err := conLdap.Search(searchReq) result, err := conLdap.Search(searchReq)
mutex.Unlock() mutex.Unlock()
if err != nil { if err != nil {
logstream.Err(fmt.Sprintln("Error searching into LDAP: ", err)) if err.Error == "ldap: connection closed" {
s := strings.Replace(err.Error(), " ", "%20", -1) // TODO : Reconnect to LDAP
response := fmt.Sprintf("500 %s\n", s) return
connClt.Write([]byte(response)) }
return logstream.Err(fmt.Sprintln("Error searching into LDAP: ", err))
} s := strings.Replace(err.Error(), " ", "%20", -1)
response := fmt.Sprintf("500 %s\n", s)
connClt.Write([]byte(response))
return
}
if len(result.Entries) != 1 { // L'attribut n'existe pas
if *debug { if len(result.Entries[0].Attributes) == 0 {
logstream.Debug("Got no result, returning 500") logstream.Err(fmt.Sprintf("Error searching into LDAP: Attribute %s not found for entry %s\n", *ldapAttr, result.Entries[0]))
} s := strings.Replace(fmt.Sprintf("Attribute not found: %s", *ldapAttr), " ", "%20", -1)
s := strings.Replace("No route defined", " ", "%20", -1) response := fmt.Sprintf("500 %s\n", s)
response := fmt.Sprintf("500 %s\n", s) connClt.Write([]byte(response))
connClt.Write([]byte(response)) return
return }
} if len(result.Entries) != 1 {
if *debug { if *debug {
logstream.Debug("Got result, returning " + result.Entries[0].Attributes[0].Values[0]) logstream.Debug("Got no result, returning 500")
} }
// Check if result is a FQDN, if not append defDomain s := strings.Replace("No route defined", " ", "%20", -1)
if strings.Contains(string(result.Entries[0].Attributes[0].Values[0]), ".") { response := fmt.Sprintf("500 %s\n", s)
s = strings.Replace(fmt.Sprintf("FILTER relay:[%s]", result.Entries[0].Attributes[0].Values[0]), " ", "%20", -1) connClt.Write([]byte(response))
} else { return
s = strings.Replace(fmt.Sprintf("FILTER relay:[%s]", fmt.Sprintf("%s.%s", result.Entries[0].Attributes[0].Values[0], *defDomain)), " ", "%20", -1) }
} if *debug {
connClt.Write([]byte(fmt.Sprintf("200 %s\n", s))) logstream.Debug("Got result, returning " + result.Entries[0].Attributes[0].Values[0])
} else { }
if *debug { // Check if result is a FQDN, if not append defDomain
logstream.Debug("Incorrect input format : " + string(buf[:readlen-1])) if strings.Contains(string(result.Entries[0].Attributes[0].Values[0]), ".") {
} s = strings.Replace(fmt.Sprintf("FILTER relay:[%s]", result.Entries[0].Attributes[0].Values[0]), " ", "%20", -1)
s := strings.Replace("Incorrect input format", " ", "%20", -1) } else {
response := fmt.Sprintf("500 %s\n", s) s = strings.Replace(fmt.Sprintf("FILTER relay:[%s]", fmt.Sprintf("%s.%s", result.Entries[0].Attributes[0].Values[0], *defDomain)), " ", "%20", -1)
connClt.Write([]byte(response)) }
return connClt.Write([]byte(fmt.Sprintf("200 %s\n", s)))
} else {
} if *debug {
logstream.Debug("Incorrect input format : " + string(buf[:readlen-1]))
}
s := strings.Replace("Incorrect input format", " ", "%20", -1)
response := fmt.Sprintf("500 %s\n", s)
connClt.Write([]byte(response))
return
}
} }
func run() { func run() {
logstream.Info("start") logstream.Info("start")
defer logstream.Info("exit") defer logstream.Info("exit")
listener, err := net.Listen("tcp", *listen) listener, err := net.Listen("tcp", *listen)
if err != nil { if err != nil {
log.Fatal(fmt.Sprintln("Error listening: ", err)) log.Fatal(fmt.Sprintln("Error listening: ", err))
} }
conLdap, err := ldap.DialURL(*ldapURL) conLdap, err := ldap.DialURL(*ldapURL)
defer conLdap.Close() defer conLdap.Close()
err = conLdap.Bind(*ldapUser, *ldapPass) err = conLdap.Bind(*ldapUser, *ldapPass)
if err != nil { if err != nil {
logstream.Err(fmt.Sprintln("Error binding LDAP: ", err)) logstream.Err(fmt.Sprintln("Error binding LDAP: ", err))
return return
} }
for { for {
connClt, err := listener.Accept() connClt, err := listener.Accept()
if err != nil { if err != nil {
logstream.Err(fmt.Sprintln("Error accepting: ", err)) logstream.Err(fmt.Sprintln("Error accepting: ", err))
} }
go handleConnection(connClt, conLdap) go handleConnection(connClt, conLdap)
} }
} }
func main() { func main() {
var e error var e error
fs := flag.NewFlagSet("mxrouter", flag.ExitOnError) fs := flag.NewFlagSet("mxrouter", flag.ExitOnError)
listen = fs.String("listen-addr", "127.0.0.1:8080", "listen address for server (also via LISTEN env var)") listen = fs.String("listen-addr", "127.0.0.1:8080", "listen address for server (also via LISTEN env var)")
debug = fs.Bool("debug", false, "log debug information (also via DEBUG env var)") debug = fs.Bool("debug", false, "log debug information (also via DEBUG env var)")
ldapURL = fs.String("ldap", "", "LDAP Server URL (also via LDAP env var)") ldapURL = fs.String("ldap", "", "LDAP Server URL (also via LDAP env var)")
ldapBaseDN = fs.String("ldapDN", "", "LDAP base DN (also via LDAPDN env var)") ldapBaseDN = fs.String("ldapDN", "", "LDAP base DN (also via LDAPDN env var)")
ldapUser = fs.String("ldapUser", "", "LDAP user DN (also via LDAPUSER env var)") ldapUser = fs.String("ldapUser", "", "LDAP user DN (also via LDAPUSER env var)")
ldapPass = fs.String("ldapPass", "", "LDAP user password (also via LDAPPASS env var)") ldapPass = fs.String("ldapPass", "", "LDAP user password (also via LDAPPASS env var)")
defDomain = fs.String("domain", "", "Domain to add to relay name if not a FQDN (also via DOMAIN env var)") ldapAttr = fs.String("ldapAttr", "relayName", "LDAP attribute containing the relay name to return")
_ = fs.String("config", "", "config file (optional)") defDomain = fs.String("domain", "", "Domain to add to relay name if not a FQDN (also via DOMAIN env var)")
_ = fs.String("config", "", "config file (optional)")
// Surcharge de la fonction Usage() // Surcharge de la fonction Usage()
fs.Usage = func() { fs.Usage = func() {
fmt.Fprintf(flag.CommandLine.Output(), "%s version %s\n", os.Args[0], version) fmt.Fprintf(flag.CommandLine.Output(), "%s version %s\n", os.Args[0], version)
fmt.Fprintf(flag.CommandLine.Output(), "Usage:\n") fmt.Fprintf(flag.CommandLine.Output(), "Usage:\n")
fs.PrintDefaults() fs.PrintDefaults()
} }
ff.Parse(fs, os.Args[1:], ff.WithEnvVarNoPrefix(), ff.WithConfigFileFlag("config"), ff.WithConfigFileParser(ff.PlainParser)) ff.Parse(fs, os.Args[1:], ff.WithEnvVarNoPrefix(), ff.WithConfigFileFlag("config"), ff.WithConfigFileParser(ff.PlainParser))
if len(*ldapURL) == 0 || len(*ldapBaseDN) == 0 || len(*ldapUser) == 0 || len(*ldapPass) == 0 { if len(*ldapURL) == 0 || len(*ldapBaseDN) == 0 || len(*ldapUser) == 0 || len(*ldapPass) == 0 {
fs.Usage() fs.Usage()
return return
} }
if logstream, e = syslog.New(syslog.LOG_MAIL, "mxrouter"); e != nil { if logstream, e = syslog.New(syslog.LOG_MAIL, "mxrouter"); e != nil {
log.Fatal(e) log.Fatal(e)
} }
defer logstream.Close() defer logstream.Close()
run() run()
} }