...

Source file src/github.com/cybertec-postgresql/pgwatch/v6/internal/sinks/rpc.go

Documentation: github.com/cybertec-postgresql/pgwatch/v6/internal/sinks

     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  // RPCWriter sends metric measurements to a remote server using gRPC.
    27  // Remote servers should make use of the .proto file under api/pb/ to integrate with it.
    28  // It's up to the implementer to define the behavior of the server.
    29  // It can be a simple logger, external storage, alerting system, or an analytics system.
    30  type RPCWriter struct {
    31  	ctx    context.Context
    32  	conn   *grpc.ClientConn
    33  	client pb.ReceiverClient
    34  }
    35  
    36  // convertSyncOp converts sinks.SyncOp to pb.SyncOp
    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  // Sends Measurement Message to RPC Sink
   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  // SyncMetric synchronizes a metric and monitored source with the remote server
   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  // DefineMetrics sends metric definitions to the remote server
   174  func (rw *RPCWriter) DefineMetrics(metrics *metrics.Metrics) error {
   175  	var json = jsoniter.ConfigFastest
   176  
   177  	// Convert metrics to JSON first, then to structpb.Struct
   178  	// to automatically handle all the type conversions
   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