diff --git a/main.go b/main.go index 774da07..172b32e 100644 --- a/main.go +++ b/main.go @@ -29,6 +29,7 @@ import ( "github.com/square/certigo/lib" "github.com/square/certigo/mysql" + "github.com/square/certigo/psql" "golang.org/x/crypto/ssh/terminal" ) @@ -50,7 +51,7 @@ var ( connectJSON = connect.Flag("json", "Write output as machine-readable JSON format.").Bool() connectCert = connect.Flag("cert", "Client certificate chain for connecting to server (PEM).").ExistingFile() connectKey = connect.Flag("key", "Private key for client certificate, if not in same file (PEM).").ExistingFile() - connectStartTLS = connect.Flag("start-tls", "Enable StartTLS protocol (supports 'mysql' for now).").PlaceHolder("PROTOCOL").Enum("mysql") + connectStartTLS = connect.Flag("start-tls", "Enable StartTLS protocol (supports 'mysql' and 'postgres').").PlaceHolder("PROTOCOL").Enum("mysql", "postgres", "psql") verify = app.Command("verify", "Verify a certificate chain from file/stdin against a name.") verifyFile = verify.Arg("file", "Certificate file to dump (or stdin if not specified).").ExistingFile() @@ -148,6 +149,9 @@ func main() { } func getConnectionState() *tls.ConnectionState { + var state *tls.ConnectionState + var err error + switch *connectStartTLS { case "": conn, err := tls.Dial("tcp", *connectTo, tlsConfigForConnect()) @@ -156,22 +160,32 @@ func getConnectionState() *tls.ConnectionState { os.Exit(1) } defer conn.Close() - state := conn.ConnectionState() - return &state + s := conn.ConnectionState() + state = &s case "mysql": mysql.RegisterTLSConfig("certigo", tlsConfigForConnect()) - state, err := mysql.DumpTLS(fmt.Sprintf("certigo@tcp(%s)/?tls=certigo", *connectTo)) - if err != nil { - fmt.Fprintf(os.Stderr, "error connecting: %v\n", err) - os.Exit(1) + state, err = mysql.DumpTLS(fmt.Sprintf("certigo@tcp(%s)/?tls=certigo", *connectTo)) + case "postgres", "psql": + // Setting sslmode to "require" skips verification. + url := fmt.Sprintf("postgres://certigo@%s/?sslmode=require", *connectTo) + if *connectCert != "" { + url += fmt.Sprintf("&sslcert=%s", *connectCert) } - - return state + if *connectKey != "" { + url += fmt.Sprintf("&sslkey=%s", *connectCert) + } + state, err = pq.DumpTLS(url) + default: + fmt.Fprintf(os.Stderr, "error connecting: unknown StartTLS protocol '%s'\n", *connectStartTLS) + os.Exit(1) } - fmt.Fprintf(os.Stderr, "error connecting: unknown StartTLS protocol '%s'\n", *connectStartTLS) - os.Exit(1) - return nil + if err != nil { + fmt.Fprintf(os.Stderr, "error connecting: %v\n", err) + os.Exit(1) + } + + return state } func inputFile(fileName string) *os.File { diff --git a/psql/conn.go b/psql/conn.go index dd91129..13eceb6 100644 --- a/psql/conn.go +++ b/psql/conn.go @@ -379,6 +379,7 @@ func DumpTLS(name string) (_ *tls.ConnectionState, err error) { defer errRecoverNoErrBadConn(&err) o := make(values) + d := defaultDialer{} // A number of defaults are applied here, in this order: // @@ -463,8 +464,13 @@ func DumpTLS(name string) (_ *tls.ConnectionState, err error) { defer cn.c.Close() if tlsConn, ok := cn.c.(*tls.Conn); ok { + err := tlsConn.Handshake() + if err != nil { + return nil, err + } + state := tlsConn.ConnectionState() - return &state + return &state, nil } return nil, errors.New("psql: TLS must be enabled to dump TLS data") diff --git a/psql/encode.go b/psql/encode.go index c3184af..50fcc55 100644 --- a/psql/encode.go +++ b/psql/encode.go @@ -13,7 +13,7 @@ import ( "sync" "time" - "github.com/square/psql/oid" + "github.com/square/certigo/psql/oid" ) func binaryEncode(parameterStatus *parameterStatus, x interface{}) []byte {