proxy-pool/internal/controller/worker/server.go
2026-08-07 18:26:09 +08:00

214 lines
6.4 KiB
Go

package worker
import (
"context"
"crypto/tls"
"errors"
"fmt"
"net"
"reflect"
"strings"
"time"
controlplanev1 "proxy-pool/gen/controlplane/v1"
"proxy-pool/internal/config"
"proxy-pool/internal/controlplane/tlsreload"
"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 {
initial, initialErr := NewInitialSnapshotSource(service, controlPlane.MaxStaleAge.Value(), time.Now)
if initialErr != nil {
return nil, fmt.Errorf("%w: build initial snapshot source: %v", ErrInvalidServer, initialErr)
}
snapshots, err = NewRefreshingSnapshotSource(initial, snapshotRefreshEvery(controlPlane.MaxStaleAge.Value()))
if err != nil {
return nil, fmt.Errorf("%w: build refreshing 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 := tlsreload.NewCertificateProvider(controlPlane.TLS.CertFile, controlPlane.TLS.KeyFile)
if err != nil {
return nil, nil, fmt.Errorf("%w: load server certificate: %v", ErrInvalidServer, err)
}
trust, err := tlsreload.NewTrustProvider(controlPlane.TLS.ClientCAFile)
if err != nil {
return nil, nil, fmt.Errorf("%w: load client ca: %v", ErrInvalidServer, err)
}
clientCAs, err := trust.Pool()
if err != nil {
return nil, nil, fmt.Errorf("%w: load client ca: %v", ErrInvalidServer, err)
}
configuration := &tls.Config{
MinVersion: tls.VersionTLS13,
GetCertificate: certificate.ServerCertificate,
ClientAuth: tls.RequireAndVerifyClientCert,
ClientCAs: clientCAs,
}
configuration.GetConfigForClient = func(*tls.ClientHelloInfo) (*tls.Config, error) {
updatedCAs, loadErr := trust.Pool()
if loadErr != nil {
return nil, loadErr
}
current := configuration.Clone()
current.GetConfigForClient = nil
current.ClientCAs = updatedCAs
return current, nil
}
return identity, []grpc.ServerOption{grpc.Creds(credentials.NewTLS(configuration))}, 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()
}