Skip to content

Commit f105c26

Browse files
committed
Add TLS and mTLS support
1 parent a3334f5 commit f105c26

6 files changed

Lines changed: 340 additions & 29 deletions

File tree

‎README.md‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -79,6 +79,29 @@ curl -o session_example.go -L https://github.com/apache/iotdb-client-go/raw/main
7979
go run session_example.go
8080
```
8181

82+
## TLS/mTLS
83+
84+
Set `TLSConfig` on `client.Config`, `client.ClusterConfig`, or `client.PoolConfig` to enable TLS. Add `CertFile` and `KeyFile` when the server requires mTLS client authentication.
85+
86+
```golang
87+
config := &client.Config{
88+
Host: host,
89+
Port: port,
90+
UserName: user,
91+
Password: password,
92+
TLSConfig: &client.TLSConfig{
93+
CAFile: "/path/to/ca.pem",
94+
CertFile: "/path/to/client.pem",
95+
KeyFile: "/path/to/client-key.pem",
96+
},
97+
}
98+
session := client.NewSession(config)
99+
if err := session.Open(false, 0); err != nil {
100+
log.Fatal(err)
101+
}
102+
defer session.Close()
103+
```
104+
82105
## How to Use the SessionPool
83106

84107
SessionPool is a wrapper of a Session Set. Using SessionPool, the user do not need to consider how to reuse a session connection.

‎README_ZH.md‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -76,6 +76,29 @@ curl -o session_example.go -L https://github.com/apache/iotdb-client-go/raw/main
7676
go run session_example.go
7777
```
7878

79+
## TLS/mTLS
80+
81+
在 `client.Config`、`client.ClusterConfig` 或 `client.PoolConfig` 上设置 `TLSConfig` 即可启用 TLS。如果服务端要求 mTLS 客户端认证,同时设置 `CertFile` 和 `KeyFile`。
82+
83+
```golang
84+
config := &client.Config{
85+
Host: host,
86+
Port: port,
87+
UserName: user,
88+
Password: password,
89+
TLSConfig: &client.TLSConfig{
90+
CAFile: "/path/to/ca.pem",
91+
CertFile: "/path/to/client.pem",
92+
KeyFile: "/path/to/client-key.pem",
93+
},
94+
}
95+
session := client.NewSession(config)
96+
if err := session.Open(false, 0); err != nil {
97+
log.Fatal(err)
98+
}
99+
defer session.Close()
100+
```
101+
79102
## SessionPool
80103
通过SessionPool管理session,用户不需要考虑如何重用session,当到达pool的最大值时,获取session的请求会阻塞
81104
注意:session使用完成后需要调用PutBack方法

‎client/session.go‎

Lines changed: 23 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,6 @@ import (
2626
"errors"
2727
"fmt"
2828
"log"
29-
"net"
3029
"reflect"
3130
"sort"
3231
"strings"
@@ -68,6 +67,7 @@ type Config struct {
6867
sqlDialect string
6968
Version Version
7069
Database string
70+
TLSConfig *TLSConfig
7171
}
7272

7373
type Session struct {
@@ -100,13 +100,10 @@ func (s *Session) Open(enableRPCCompression bool, connectionTimeoutInMs int) err
100100

101101
var err error
102102

103-
// In thrift 0.14.1, this func returns two values; in newer versions, it returns one.
104-
s.trans = thrift.NewTSocketConf(net.JoinHostPort(s.config.Host, s.config.Port), &thrift.TConfiguration{
105-
ConnectTimeout: time.Duration(connectionTimeoutInMs) * time.Millisecond, // Use 0 for no timeout
106-
})
107-
// s.trans = thrift.NewTFramedTransport(s.trans) // deprecated
108-
tmp_conf := thrift.TConfiguration{MaxFrameSize: thrift.DEFAULT_MAX_FRAME_SIZE}
109-
s.trans = thrift.NewTFramedTransportConf(s.trans, &tmp_conf)
103+
s.trans, err = newTransport(s.config.Host, s.config.Port, connectionTimeoutInMs, s.config.TLSConfig)
104+
if err != nil {
105+
return err
106+
}
110107
if !s.trans.IsOpen() {
111108
err = s.trans.Open()
112109
if err != nil {
@@ -154,6 +151,7 @@ type ClusterConfig struct {
154151
ConnectRetryMax int
155152
sqlDialect string
156153
Database string
154+
TLSConfig *TLSConfig
157155
}
158156

159157
func (s *Session) OpenCluster(enableRPCCompression bool) error {
@@ -1328,24 +1326,23 @@ func newClusterSessionWithSqlDialect(clusterConfig *ClusterConfig) (Session, err
13281326
var err error
13291327
for i := range session.endPointList {
13301328
ep := session.endPointList[i]
1331-
session.trans = thrift.NewTSocketConf(net.JoinHostPort(ep.Host, ep.Port), &thrift.TConfiguration{
1332-
ConnectTimeout: time.Duration(0), // Use 0 for no timeout
1333-
})
1334-
// session.trans = thrift.NewTFramedTransport(session.trans) // deprecated
1335-
tmp_conf := thrift.TConfiguration{MaxFrameSize: thrift.DEFAULT_MAX_FRAME_SIZE}
1336-
session.trans = thrift.NewTFramedTransportConf(session.trans, &tmp_conf)
1329+
session.trans, err = newTransport(ep.Host, ep.Port, 0, clusterConfig.TLSConfig)
1330+
if err != nil {
1331+
log.Println(err)
1332+
continue
1333+
}
13371334
if !session.trans.IsOpen() {
13381335
err = session.trans.Open()
13391336
if err != nil {
13401337
log.Println(err)
13411338
} else {
13421339
session.config = getConfig(ep.Host, ep.Port,
1343-
clusterConfig.UserName, clusterConfig.Password, clusterConfig.FetchSize, clusterConfig.TimeZone, clusterConfig.ConnectRetryMax, clusterConfig.Database, clusterConfig.sqlDialect)
1340+
clusterConfig.UserName, clusterConfig.Password, clusterConfig.FetchSize, clusterConfig.TimeZone, clusterConfig.ConnectRetryMax, clusterConfig.Database, clusterConfig.sqlDialect, clusterConfig.TLSConfig)
13441341
break
13451342
}
13461343
}
13471344
}
1348-
if !session.trans.IsOpen() {
1345+
if session.trans == nil || !session.trans.IsOpen() {
13491346
return session, fmt.Errorf("no server can connect")
13501347
}
13511348
return session, nil
@@ -1354,18 +1351,14 @@ func newClusterSessionWithSqlDialect(clusterConfig *ClusterConfig) (Session, err
13541351
func (s *Session) initClusterConn(node endPoint) error {
13551352
var err error
13561353

1357-
s.trans = thrift.NewTSocketConf(net.JoinHostPort(node.Host, node.Port), &thrift.TConfiguration{
1358-
ConnectTimeout: time.Duration(0), // Use 0 for no timeout
1359-
})
1360-
if err == nil {
1361-
// s.trans = thrift.NewTFramedTransport(s.trans) // deprecated
1362-
tmp_conf := thrift.TConfiguration{MaxFrameSize: thrift.DEFAULT_MAX_FRAME_SIZE}
1363-
s.trans = thrift.NewTFramedTransportConf(s.trans, &tmp_conf)
1364-
if !s.trans.IsOpen() {
1365-
err = s.trans.Open()
1366-
if err != nil {
1367-
return err
1368-
}
1354+
s.trans, err = newTransport(node.Host, node.Port, 0, s.config.TLSConfig)
1355+
if err != nil {
1356+
return err
1357+
}
1358+
if !s.trans.IsOpen() {
1359+
err = s.trans.Open()
1360+
if err != nil {
1361+
return err
13691362
}
13701363
}
13711364

@@ -1398,7 +1391,7 @@ func (s *Session) initClusterConn(node endPoint) error {
13981391
return err
13991392
}
14001393

1401-
func getConfig(host string, port string, userName string, passWord string, fetchSize int32, timeZone string, connectRetryMax int, database string, sqlDialect string) *Config {
1394+
func getConfig(host string, port string, userName string, passWord string, fetchSize int32, timeZone string, connectRetryMax int, database string, sqlDialect string, tlsConfig *TLSConfig) *Config {
14021395
return &Config{
14031396
Host: host,
14041397
Port: port,
@@ -1409,6 +1402,7 @@ func getConfig(host string, port string, userName string, passWord string, fetch
14091402
ConnectRetryMax: connectRetryMax,
14101403
sqlDialect: sqlDialect,
14111404
Database: database,
1405+
TLSConfig: tlsConfig,
14121406
}
14131407
}
14141408

‎client/sessionpool.go‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,7 @@ type PoolConfig struct {
5050
TimeZone string
5151
ConnectRetryMax int
5252
Database string
53+
TLSConfig *TLSConfig
5354
sqlDialect string
5455
}
5556

@@ -146,6 +147,7 @@ func getSessionConfig(config *PoolConfig) *Config {
146147
ConnectRetryMax: config.ConnectRetryMax,
147148
sqlDialect: config.sqlDialect,
148149
Database: config.Database,
150+
TLSConfig: config.TLSConfig,
149151
}
150152
}
151153

@@ -159,6 +161,7 @@ func getClusterSessionConfig(config *PoolConfig) *ClusterConfig {
159161
ConnectRetryMax: config.ConnectRetryMax,
160162
sqlDialect: config.sqlDialect,
161163
Database: config.Database,
164+
TLSConfig: config.TLSConfig,
162165
}
163166
}
164167

‎client/tls.go‎

Lines changed: 120 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,120 @@
1+
/*
2+
* Licensed to the Apache Software Foundation (ASF) under one
3+
* or more contributor license agreements. See the NOTICE file
4+
* distributed with this work for additional information
5+
* regarding copyright ownership. The ASF licenses this file
6+
* to you under the Apache License, Version 2.0 (the
7+
* "License"); you may not use this file except in compliance
8+
* with the License. You may obtain a copy of the License at
9+
*
10+
* http://www.apache.org/licenses/LICENSE-2.0
11+
*
12+
* Unless required by applicable law or agreed to in writing,
13+
* software distributed under the License is distributed on an
14+
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15+
* KIND, either express or implied. See the License for the
16+
* specific language governing permissions and limitations
17+
* under the License.
18+
*/
19+
20+
package client
21+
22+
import (
23+
"crypto/tls"
24+
"crypto/x509"
25+
"fmt"
26+
"net"
27+
"os"
28+
"time"
29+
30+
"github.com/apache/thrift/lib/go/thrift"
31+
)
32+
33+
// TLSConfig enables TLS for an IoTDB client connection. Set CertFile and
34+
// KeyFile together to enable mTLS client authentication.
35+
type TLSConfig struct {
36+
// Config is an optional base tls.Config. It is cloned before use.
37+
Config *tls.Config
38+
39+
// CAFile is an optional PEM encoded CA certificate file used to verify the server.
40+
CAFile string
41+
42+
// CertFile and KeyFile are optional PEM encoded client certificate and key files for mTLS.
43+
CertFile string
44+
KeyFile string
45+
}
46+
47+
func newTransport(host string, port string, connectionTimeoutInMs int, tlsConfig *TLSConfig) (thrift.TTransport, error) {
48+
conf := &thrift.TConfiguration{
49+
ConnectTimeout: time.Duration(connectionTimeoutInMs) * time.Millisecond,
50+
MaxFrameSize: thrift.DEFAULT_MAX_FRAME_SIZE,
51+
}
52+
hostPort := net.JoinHostPort(host, port)
53+
54+
var base thrift.TTransport
55+
if tlsConfig == nil {
56+
base = thrift.NewTSocketConf(hostPort, conf)
57+
} else {
58+
cfg, err := buildTLSConfig(tlsConfig)
59+
if err != nil {
60+
return nil, err
61+
}
62+
conf.TLSConfig = cfg
63+
base = thrift.NewTSSLSocketConf(hostPort, conf)
64+
}
65+
66+
return thrift.NewTFramedTransportConf(base, conf), nil
67+
}
68+
69+
func buildTLSConfig(config *TLSConfig) (*tls.Config, error) {
70+
if config == nil {
71+
return nil, nil
72+
}
73+
74+
tlsConfig := &tls.Config{}
75+
if config.Config != nil {
76+
tlsConfig = config.Config.Clone()
77+
}
78+
if config.CAFile != "" {
79+
rootCAs, err := loadCertPool(tlsConfig.RootCAs, config.CAFile)
80+
if err != nil {
81+
return nil, err
82+
}
83+
tlsConfig.RootCAs = rootCAs
84+
}
85+
if config.CertFile != "" || config.KeyFile != "" {
86+
if config.CertFile == "" || config.KeyFile == "" {
87+
return nil, fmt.Errorf("both TLS CertFile and KeyFile must be set")
88+
}
89+
certificate, err := tls.LoadX509KeyPair(config.CertFile, config.KeyFile)
90+
if err != nil {
91+
return nil, fmt.Errorf("load TLS client certificate/key: %w", err)
92+
}
93+
tlsConfig.Certificates = append(tlsConfig.Certificates, certificate)
94+
}
95+
96+
return tlsConfig, nil
97+
}
98+
99+
func loadCertPool(base *x509.CertPool, caFile string) (*x509.CertPool, error) {
100+
rootCAs := base
101+
if rootCAs != nil {
102+
rootCAs = rootCAs.Clone()
103+
} else {
104+
systemPool, err := x509.SystemCertPool()
105+
if err == nil {
106+
rootCAs = systemPool
107+
} else {
108+
rootCAs = x509.NewCertPool()
109+
}
110+
}
111+
112+
caCert, err := os.ReadFile(caFile)
113+
if err != nil {
114+
return nil, fmt.Errorf("read TLS CA file %q: %w", caFile, err)
115+
}
116+
if !rootCAs.AppendCertsFromPEM(caCert) {
117+
return nil, fmt.Errorf("append TLS CA file %q: no certificates found", caFile)
118+
}
119+
return rootCAs, nil
120+
}

0 commit comments

Comments
 (0)