1 package sinks
2
3 import (
4 "context"
5 "crypto/tls"
6 "crypto/x509"
7 "errors"
8 "fmt"
9 "net/url"
10 "os"
11 "time"
12
13 "github.com/cybertec-postgresql/pgwatch/v6/api/pb"
14 "github.com/cybertec-postgresql/pgwatch/v6/internal/log"
15 "github.com/cybertec-postgresql/pgwatch/v6/internal/metrics"
16 jsoniter "github.com/json-iterator/go"
17 "google.golang.org/grpc"
18 "google.golang.org/grpc/codes"
19 "google.golang.org/grpc/credentials"
20 "google.golang.org/grpc/credentials/insecure"
21 "google.golang.org/grpc/metadata"
22 "google.golang.org/grpc/status"
23 "google.golang.org/protobuf/types/known/structpb"
24 )
25
26
27
28
29
30 type RPCWriter struct {
31 ctx context.Context
32 conn *grpc.ClientConn
33 client pb.ReceiverClient
34 }
35
36
37 func convertSyncOp(op SyncOp) pb.SyncOp {
38 switch op {
39 case AddOp:
40 return pb.SyncOp_AddOp
41 case DeleteOp:
42 return pb.SyncOp_DeleteOp
43 case DefineOp:
44 return pb.SyncOp_DefineOp
45 default:
46 return pb.SyncOp_InvalidOp
47 }
48 }
49
50 func NewRPCWriter(ctx context.Context, connStr string) (*RPCWriter, error) {
51 uri, err := url.Parse(connStr)
52 if err != nil {
53 return nil, fmt.Errorf("error parsing gRPC URI: %s", err)
54 }
55
56 l := log.GetLogger(ctx).WithField("sink", "grpc").WithField("address", uri.Host)
57 ctx = log.WithLogger(ctx, l)
58
59 params, err := url.ParseQuery(uri.RawQuery)
60 if err != nil {
61 return nil, fmt.Errorf("error parsing gRPC URI parameters: %s", err)
62 }
63
64 creds := insecure.NewCredentials()
65
66 CAFile, ok := params["sslrootca"]
67 if ok {
68 creds, err = LoadTLSCredentials(CAFile[0])
69 if err != nil {
70 return nil, err
71 }
72 log.GetLogger(ctx).Infof("Valid CA File %s loaded - enabling TLS", CAFile)
73 }
74
75 conn, err := grpc.NewClient(uri.Host, grpc.WithTransportCredentials(creds))
76 if err != nil {
77 return nil, err
78 }
79
80 password, _ := uri.User.Password()
81 md := metadata.Pairs(
82 "username", uri.User.Username(),
83 "password", password,
84 )
85 newCtx := metadata.NewOutgoingContext(ctx, md)
86
87 client := pb.NewReceiverClient(conn)
88 rw := &RPCWriter{
89 ctx: newCtx,
90 conn: conn,
91 client: client,
92 }
93
94 if err = rw.Ping(); err != nil {
95 return nil, err
96 }
97
98 go rw.watchCtx()
99 return rw, nil
100 }
101
102 func (rw *RPCWriter) Ping() error {
103 err := rw.SyncMetric("", "", InvalidOp)
104 st, ok := status.FromError(err)
105 if ok && st.Code() == codes.Unavailable {
106 return err
107 }
108 return nil
109 }
110
111
112 func (rw *RPCWriter) Write(msg metrics.MeasurementEnvelope) error {
113 if rw.ctx.Err() != nil {
114 return rw.ctx.Err()
115 }
116
117 dataLength := len(msg.Data)
118 failCnt := 0
119 measurements := make([]*structpb.Struct, 0, dataLength)
120 for _, item := range msg.Data {
121 st, err := structpb.NewStruct(item)
122 if err != nil {
123 failCnt++
124 continue
125 }
126 measurements = append(measurements, st)
127 }
128 if failCnt > 0 {
129 log.GetLogger(rw.ctx).WithField("database", msg.DBName).WithField("metric",
130 msg.MetricName).Warningf("gRPC sink failed to encode %d rows", failCnt)
131 }
132
133 envelope := &pb.MeasurementEnvelope{
134 DBName: msg.DBName,
135 MetricName: msg.MetricName,
136 CustomTags: msg.CustomTags,
137 Data: measurements,
138 }
139
140 t1 := time.Now()
141 reply, err := rw.client.UpdateMeasurements(rw.ctx, envelope)
142 if err != nil {
143 return err
144 }
145
146 diff := time.Since(t1)
147 log.GetLogger(rw.ctx).WithField("rows", dataLength).WithField("elapsed", diff).Info("measurements written")
148 if reply.GetLogmsg() != "" {
149 log.GetLogger(rw.ctx).Info(reply.GetLogmsg())
150 }
151 return nil
152 }
153
154
155 func (rw *RPCWriter) SyncMetric(sourceName, metricName string, op SyncOp) error {
156 syncReq := &pb.SyncReq{
157 DBName: sourceName,
158 MetricName: metricName,
159 Operation: convertSyncOp(op),
160 }
161
162 reply, err := rw.client.SyncMetric(rw.ctx, syncReq)
163 if err != nil {
164 return err
165 }
166
167 if reply.GetLogmsg() != "" {
168 log.GetLogger(rw.ctx).Info(reply.GetLogmsg())
169 }
170 return nil
171 }
172
173
174 func (rw *RPCWriter) DefineMetrics(metrics *metrics.Metrics) error {
175 var json = jsoniter.ConfigFastest
176
177
178
179 jsonData, err := json.Marshal(metrics)
180 if err != nil {
181 return err
182 }
183
184 var metricMap map[string]any
185 if err := json.Unmarshal(jsonData, &metricMap); err != nil {
186 return err
187 }
188
189 metricStruct, err := structpb.NewStruct(metricMap)
190 if err != nil {
191 return err
192 }
193
194 t1 := time.Now()
195 reply, err := rw.client.DefineMetrics(rw.ctx, metricStruct)
196 if err != nil {
197 return err
198 }
199
200 diff := time.Since(t1)
201 log.GetLogger(rw.ctx).WithField("elapsed", diff).Info("metric definitions written")
202 if reply.GetLogmsg() != "" {
203 log.GetLogger(rw.ctx).Info(reply.GetLogmsg())
204 }
205 return nil
206 }
207
208 func (rw *RPCWriter) watchCtx() {
209 <-rw.ctx.Done()
210 rw.conn.Close()
211 }
212
213 func LoadTLSCredentials(CAFile string) (credentials.TransportCredentials, error) {
214 ca, err := os.ReadFile(CAFile)
215 if err != nil {
216 return nil, fmt.Errorf("error loading CA file: %v", err)
217 }
218
219 certPool := x509.NewCertPool()
220 ok := certPool.AppendCertsFromPEM(ca)
221 if !ok {
222 return nil, errors.New("invalid CA file")
223 }
224
225 tlsClientConfig := &tls.Config{
226 RootCAs: certPool,
227 }
228 return credentials.NewTLS(tlsClientConfig), nil
229 }
230