mirror of
https://github.com/square/certigo.git
synced 2024-04-21 12:32:40 +00:00
Merge pull request #157 from square/mmc/default-port
Add protocol-specific default ports
This commit is contained in:
@@ -44,7 +44,7 @@ var (
|
||||
dumpJSON = dump.Flag("json", "Write output as machine-readable JSON format.").Short('j').Bool()
|
||||
|
||||
connect = app.Command("connect", "Connect to a server and print its certificate(s).")
|
||||
connectTo = connect.Arg("server:port", "Hostname or IP to connect to.").String()
|
||||
connectTo = connect.Arg("server[:port]", "Hostname or IP to connect to, with optional port.").String()
|
||||
connectName = connect.Flag("name", "Override the server name used for Server Name Indication (SNI).").Short('n').String()
|
||||
connectCaPath = connect.Flag("ca", "Path to CA bundle (system default if unspecified).").ExistingFile()
|
||||
connectCert = connect.Flag("cert", "Client certificate chain for connecting to server (PEM).").ExistingFile()
|
||||
|
||||
+32
-6
@@ -23,6 +23,8 @@ import (
|
||||
"net"
|
||||
"net/smtp"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/square/certigo/starttls/ldap"
|
||||
@@ -78,6 +80,24 @@ func setGetClientCertificateCallback(conf *tls.Config, cert *tls.Certificate) **
|
||||
return &captured
|
||||
}
|
||||
|
||||
// withDefaultPort takes an address and a port, and returns the address
|
||||
// as-is if there's a port, or address:port if there isn't.
|
||||
func withDefaultPort(addr string, portN uint16) string {
|
||||
port := strconv.Itoa(int(portN))
|
||||
if strings.HasPrefix(addr, "[") && strings.HasSuffix(addr, "]") {
|
||||
// IPv6 without a port
|
||||
return net.JoinHostPort(addr, port)
|
||||
}
|
||||
if strings.ContainsAny(addr, ":") {
|
||||
// There are two cases when an address could contain a `:`
|
||||
// If it's an IPv6 address without a port, it's handled above
|
||||
// Otherwise, it's any kind of address with a port
|
||||
return addr
|
||||
}
|
||||
// No colon means no port.
|
||||
return net.JoinHostPort(addr, port)
|
||||
}
|
||||
|
||||
// GetConnectionState connects to a TLS server, returning the connection state.
|
||||
// Currently, startTLSType can be one of "mysql", "postgres" or "psql", or the
|
||||
// empty string, which does a normal TLS connection. connectTo specifies the
|
||||
@@ -121,7 +141,8 @@ func GetConnectionState(startTLSType, connectName, connectTo, identity, clientCe
|
||||
go func() {
|
||||
switch startTLSType {
|
||||
case "":
|
||||
conn, err := dialWithDialer(dialer, timeout, "tcp", connectTo, tlsConfig)
|
||||
addr := withDefaultPort(connectTo, 443)
|
||||
conn, err := dialWithDialer(dialer, timeout, "tcp", addr, tlsConfig)
|
||||
if err != nil {
|
||||
res <- connectResult{nil, err}
|
||||
return
|
||||
@@ -130,7 +151,8 @@ func GetConnectionState(startTLSType, connectName, connectTo, identity, clientCe
|
||||
state := conn.ConnectionState()
|
||||
res <- connectResult{&state, nil}
|
||||
case "ldap":
|
||||
l, err := ldap.Dial("tcp", connectTo, timeout)
|
||||
addr := withDefaultPort(connectTo, 389)
|
||||
l, err := ldap.Dial("tcp", addr, timeout)
|
||||
if err != nil {
|
||||
res <- connectResult{nil, err}
|
||||
return
|
||||
@@ -150,7 +172,8 @@ func GetConnectionState(startTLSType, connectName, connectTo, identity, clientCe
|
||||
res <- connectResult{state, nil}
|
||||
case "mysql":
|
||||
mysql.RegisterTLSConfig("certigo", tlsConfig)
|
||||
state, err = mysql.DumpTLS(fmt.Sprintf("%s@tcp(%s)/?tls=certigo&timeout=%s", identity, connectTo, timeout.String()))
|
||||
addr := withDefaultPort(connectTo, 3306)
|
||||
state, err = mysql.DumpTLS(fmt.Sprintf("%s@tcp(%s)/?tls=certigo&timeout=%s", identity, addr, timeout.String()))
|
||||
if err != nil {
|
||||
res <- connectResult{nil, err}
|
||||
return
|
||||
@@ -158,7 +181,8 @@ func GetConnectionState(startTLSType, connectName, connectTo, identity, clientCe
|
||||
res <- connectResult{state, nil}
|
||||
case "postgres", "psql":
|
||||
// Setting sslmode to "require" skips verification.
|
||||
url := fmt.Sprintf("postgres://%s@%s/?sslmode=require&connect_timeout=%d", identity, connectTo, timeout/time.Second)
|
||||
addr := withDefaultPort(connectTo, 5432)
|
||||
url := fmt.Sprintf("postgres://%s@%s/?sslmode=require&connect_timeout=%d", identity, addr, timeout/time.Second)
|
||||
if clientCert != "" {
|
||||
url += fmt.Sprintf("&sslcert=%s", clientCert)
|
||||
}
|
||||
@@ -175,7 +199,8 @@ func GetConnectionState(startTLSType, connectName, connectTo, identity, clientCe
|
||||
// Go's net/smtp doesn't support timeouts, so if we hit a timeout we might
|
||||
// leak a Go routine (at least until we hit a lower-level TCP timeout or such).
|
||||
// This is not an issue for Certigo since it's just a short-lived CLI utility.
|
||||
client, err := smtp.Dial(connectTo)
|
||||
addr := withDefaultPort(connectTo, 25)
|
||||
client, err := smtp.Dial(addr)
|
||||
if err != nil {
|
||||
res <- connectResult{nil, err}
|
||||
return
|
||||
@@ -196,7 +221,8 @@ func GetConnectionState(startTLSType, connectName, connectTo, identity, clientCe
|
||||
}
|
||||
res <- connectResult{&state, nil}
|
||||
case "ftp":
|
||||
state, err = dumpAuthTLSFromFTP(dialer, connectTo, tlsConfig)
|
||||
addr := withDefaultPort(connectTo, 21)
|
||||
state, err = dumpAuthTLSFromFTP(dialer, addr, tlsConfig)
|
||||
res <- connectResult{state, err}
|
||||
default:
|
||||
res <- connectResult{nil, fmt.Errorf("unknown StartTLS protocol: %s", startTLSType)}
|
||||
|
||||
Reference in New Issue
Block a user