201 lines
5.8 KiB
Go
201 lines
5.8 KiB
Go
package worker
|
|
|
|
import (
|
|
"context"
|
|
"crypto/tls"
|
|
"crypto/x509"
|
|
"errors"
|
|
"fmt"
|
|
"net"
|
|
"os"
|
|
"reflect"
|
|
"strings"
|
|
"time"
|
|
|
|
controlplanev1 "proxy-pool/gen/controlplane/v1"
|
|
"proxy-pool/internal/config"
|
|
|
|
"google.golang.org/grpc"
|
|
"google.golang.org/grpc/credentials"
|
|
"google.golang.org/grpc/keepalive"
|
|
)
|
|
|
|
var ErrInvalidServer = errors.New("invalid worker control server configuration")
|
|
|
|
type ServerOptions struct {
|
|
ShutdownTimeout time.Duration
|
|
Snapshots SnapshotSource
|
|
Checker controlplanev1.CheckerControlPlaneServer
|
|
}
|
|
|
|
func DefaultServerOptions() ServerOptions {
|
|
return ServerOptions{ShutdownTimeout: 15 * time.Second}
|
|
}
|
|
|
|
type Server struct {
|
|
listen string
|
|
grpcServer *grpc.Server
|
|
shutdownTimeout time.Duration
|
|
}
|
|
|
|
func NewServer(controlPlane config.ControlPlane, service Service, options ServerOptions) (*Server, error) {
|
|
if service == nil || !controlPlane.Enabled || !validServerConfig(controlPlane) {
|
|
return nil, ErrInvalidServer
|
|
}
|
|
if options.ShutdownTimeout < 0 {
|
|
return nil, fmt.Errorf("%w: shutdown timeout must not be negative", ErrInvalidServer)
|
|
}
|
|
if options.ShutdownTimeout == 0 {
|
|
options = DefaultServerOptions()
|
|
}
|
|
|
|
identity, serverOptions, err := serverTransportOptions(controlPlane)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
snapshots := options.Snapshots
|
|
if snapshots == nil {
|
|
if provider, ok := service.(interface{ SnapshotSource() SnapshotSource }); ok {
|
|
snapshots = provider.SnapshotSource()
|
|
}
|
|
}
|
|
if snapshots == nil {
|
|
snapshots, err = NewInitialSnapshotSource(service, controlPlane.MaxStaleAge.Value(), time.Now)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("%w: build initial snapshot source: %v", ErrInvalidServer, err)
|
|
}
|
|
}
|
|
serverOptions = append(serverOptions,
|
|
grpc.MaxRecvMsgSize(controlPlane.MaxMessageBytes),
|
|
grpc.MaxSendMsgSize(controlPlane.MaxMessageBytes),
|
|
grpc.MaxConcurrentStreams(controlPlane.MaxConcurrentStreams),
|
|
grpc.KeepaliveEnforcementPolicy(keepalive.EnforcementPolicy{
|
|
MinTime: 10 * time.Second,
|
|
PermitWithoutStream: false,
|
|
}),
|
|
)
|
|
grpcServer := grpc.NewServer(serverOptions...)
|
|
controlplanev1.RegisterWorkerControlPlaneServer(grpcServer, NewGRPCHandler(service, identity, snapshots))
|
|
if !nilService(options.Checker) {
|
|
controlplanev1.RegisterCheckerControlPlaneServer(grpcServer, options.Checker)
|
|
}
|
|
return &Server{listen: controlPlane.Listen, grpcServer: grpcServer, shutdownTimeout: options.ShutdownTimeout}, nil
|
|
}
|
|
|
|
func nilService(value any) bool {
|
|
if value == nil {
|
|
return true
|
|
}
|
|
reflected := reflect.ValueOf(value)
|
|
switch reflected.Kind() {
|
|
case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice:
|
|
return reflected.IsNil()
|
|
default:
|
|
return false
|
|
}
|
|
}
|
|
|
|
func (server *Server) Run(ctx context.Context) error {
|
|
if server == nil || server.grpcServer == nil || server.listen == "" {
|
|
return ErrInvalidServer
|
|
}
|
|
listener, err := net.Listen("tcp", server.listen)
|
|
if err != nil {
|
|
return fmt.Errorf("listen worker control plane: %w", err)
|
|
}
|
|
return server.Serve(ctx, listener)
|
|
}
|
|
|
|
func (server *Server) Serve(ctx context.Context, listener net.Listener) error {
|
|
if server == nil || server.grpcServer == nil || listener == nil || ctx == nil {
|
|
return ErrInvalidServer
|
|
}
|
|
completed := make(chan struct{})
|
|
go func() {
|
|
select {
|
|
case <-ctx.Done():
|
|
server.gracefulStop()
|
|
case <-completed:
|
|
}
|
|
}()
|
|
|
|
err := server.grpcServer.Serve(listener)
|
|
close(completed)
|
|
if ctx.Err() != nil || errors.Is(err, grpc.ErrServerStopped) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
|
|
func (server *Server) gracefulStop() {
|
|
stopped := make(chan struct{})
|
|
go func() {
|
|
server.grpcServer.GracefulStop()
|
|
close(stopped)
|
|
}()
|
|
timer := time.NewTimer(server.shutdownTimeout)
|
|
defer timer.Stop()
|
|
select {
|
|
case <-stopped:
|
|
case <-timer.C:
|
|
server.grpcServer.Stop()
|
|
<-stopped
|
|
}
|
|
}
|
|
|
|
func serverTransportOptions(controlPlane config.ControlPlane) (IdentityAuthorizer, []grpc.ServerOption, error) {
|
|
switch controlPlane.TLS.Mode {
|
|
case "disabled":
|
|
if !loopbackListen(controlPlane.Listen) {
|
|
return nil, nil, fmt.Errorf("%w: plaintext listener must be loopback", ErrInvalidServer)
|
|
}
|
|
return AllowLoopbackIdentity{}, nil, nil
|
|
case "mtls":
|
|
identity, err := NewSPIFFEIdentityAuthorizer(controlPlane.TLS.TrustDomain, controlPlane.TLS.Environment)
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("%w: %v", ErrInvalidServer, err)
|
|
}
|
|
certificate, err := tls.LoadX509KeyPair(controlPlane.TLS.CertFile, controlPlane.TLS.KeyFile)
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("%w: load server certificate: %v", ErrInvalidServer, err)
|
|
}
|
|
caPEM, err := os.ReadFile(controlPlane.TLS.ClientCAFile)
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("%w: read client ca: %v", ErrInvalidServer, err)
|
|
}
|
|
clientCAs := x509.NewCertPool()
|
|
if !clientCAs.AppendCertsFromPEM(caPEM) {
|
|
return nil, nil, fmt.Errorf("%w: parse client ca", ErrInvalidServer)
|
|
}
|
|
transport := credentials.NewTLS(&tls.Config{
|
|
MinVersion: tls.VersionTLS13,
|
|
Certificates: []tls.Certificate{certificate},
|
|
ClientAuth: tls.RequireAndVerifyClientCert,
|
|
ClientCAs: clientCAs,
|
|
})
|
|
return identity, []grpc.ServerOption{grpc.Creds(transport)}, nil
|
|
default:
|
|
return nil, nil, fmt.Errorf("%w: unsupported tls mode", ErrInvalidServer)
|
|
}
|
|
}
|
|
|
|
func validServerConfig(controlPlane config.ControlPlane) bool {
|
|
return controlPlane.Listen != "" && controlPlane.ProtocolVersion == 1 &&
|
|
controlPlane.HeartbeatInterval.Value() > 0 && controlPlane.SessionTTL.Value() > 0 &&
|
|
controlPlane.MaxStaleAge.Value() > 0 && controlPlane.MaxMessageBytes > 0 &&
|
|
controlPlane.MaxRuntimeCounters > 0 && controlPlane.MaxConcurrentStreams > 0
|
|
}
|
|
|
|
func loopbackListen(listen string) bool {
|
|
host, _, err := net.SplitHostPort(listen)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
host = strings.Trim(host, "[]")
|
|
if strings.EqualFold(host, "localhost") {
|
|
return true
|
|
}
|
|
ip := net.ParseIP(host)
|
|
return ip != nil && ip.IsLoopback()
|
|
}
|