main
go 145 lines 3.65 KB
Raw
1 package utils
2
3 import (
4 "context"
5 "crypto"
6 "crypto/ecdsa"
7 "crypto/rsa"
8 "crypto/tls"
9 "crypto/x509"
10 "encoding/pem"
11 "errors"
12 "fmt"
13 "net"
14 "net/http"
15 "net/url"
16 "strings"
17 "time"
18 )
19
20 func NewHTTPTLSClient(ctx context.Context, relayURL *url.URL, timeout time.Duration) (*tls.Config, *http.Client, *http.Transport, error) {
21 if relayURL == nil {
22 return nil, nil, nil, errors.New("relay url is required")
23 }
24
25 serverName := relayURL.Hostname()
26 if serverName == "" {
27 return nil, nil, nil, errors.New("relay hostname is required")
28 }
29
30 var rootCAs *x509.CertPool
31 if IsLocalRelayHost(serverName) {
32 rootCAPEM, err := FetchEndpointCertificateChain(ctx, relayURL.String(), serverName)
33 if err != nil {
34 return nil, nil, nil, fmt.Errorf("bootstrap localhost relay trust: %w", err)
35 }
36 rootCAs = x509.NewCertPool()
37 if !rootCAs.AppendCertsFromPEM(rootCAPEM) {
38 return nil, nil, nil, errors.New("failed to parse relay root ca")
39 }
40 }
41
42 rawTLSConfig := &tls.Config{
43 MinVersion: tls.VersionTLS12,
44 ServerName: serverName,
45 RootCAs: rootCAs,
46 NextProtos: []string{"http/1.1"},
47 }
48 httpClient := NewHTTPClient(
49 WithHTTPTLSConfig(rawTLSConfig), // will be cloned internally
50 WithoutHTTP2(),
51 WithHTTPTimeout(timeout),
52 )
53 return rawTLSConfig, httpClient, mustTransportOf(httpClient), nil
54 }
55
56 func FetchEndpointCertificateChain(ctx context.Context, endpoint, serverName string) ([]byte, error) {
57 raw := strings.TrimSpace(endpoint)
58 if raw == "" {
59 return nil, errors.New("endpoint is required")
60 }
61 if !strings.Contains(raw, "://") {
62 raw = "https://" + raw
63 }
64
65 u, err := url.Parse(raw)
66 if err != nil {
67 return nil, fmt.Errorf("parse endpoint url: %w", err)
68 }
69 if !strings.EqualFold(u.Scheme, "https") {
70 return nil, errors.New("relay endpoint must use https")
71 }
72
73 host := u.Hostname()
74 if host == "" {
75 return nil, errors.New("endpoint hostname is empty")
76 }
77 port := u.Port()
78 if port == "" {
79 port = "443"
80 }
81 if serverName == "" {
82 serverName = host
83 }
84
85 dialer := &net.Dialer{Timeout: 5 * time.Second}
86 rawConn, err := dialer.DialContext(ctx, "tcp", net.JoinHostPort(host, port))
87 if err != nil {
88 return nil, fmt.Errorf("dial relay endpoint: %w", err)
89 }
90
91 tlsConn := tls.Client(rawConn, &tls.Config{
92 MinVersion: tls.VersionTLS12,
93 ServerName: serverName,
94 InsecureSkipVerify: IsLocalRelayHost(host),
95 NextProtos: []string{"http/1.1"},
96 })
97 defer tlsConn.Close()
98 if err := tlsConn.HandshakeContext(ctx); err != nil {
99 return nil, fmt.Errorf("tls handshake with relay endpoint: %w", err)
100 }
101
102 peerCerts := tlsConn.ConnectionState().PeerCertificates
103 if len(peerCerts) == 0 {
104 return nil, errors.New("no peer certificates from relay endpoint")
105 }
106
107 var chainPEM []byte
108 for _, cert := range peerCerts {
109 chainPEM = append(chainPEM, pem.EncodeToMemory(&pem.Block{
110 Type: "CERTIFICATE",
111 Bytes: cert.Raw,
112 })...)
113 }
114 return chainPEM, nil
115 }
116
117 func ParseCertificatePEM(pemData []byte) (*x509.Certificate, error) {
118 block, _ := pem.Decode(pemData)
119 if block == nil {
120 return nil, errors.New("no pem block found")
121 }
122 return x509.ParseCertificate(block.Bytes)
123 }
124
125 func ParsePrivateKeyPEM(keyPEM []byte) (crypto.PrivateKey, error) {
126 block, _ := pem.Decode(keyPEM)
127 if block == nil {
128 return nil, errors.New("invalid private key pem")
129 }
130 if key, err := x509.ParsePKCS8PrivateKey(block.Bytes); err == nil {
131 switch typed := key.(type) {
132 case *ecdsa.PrivateKey:
133 return typed, nil
134 case *rsa.PrivateKey:
135 return typed, nil
136 }
137 }
138 if key, err := x509.ParseECPrivateKey(block.Bytes); err == nil {
139 return key, nil
140 }
141 if key, err := x509.ParsePKCS1PrivateKey(block.Bytes); err == nil {
142 return key, nil
143 }
144 return nil, errors.New("unsupported private key type")
145 }