proxy-pool/internal/controller/worker/server.go
youfak 6b6fb54075
Some checks are pending
ci / proto (push) Waiting to run
ci / test (ubuntu-latest) (push) Waiting to run
ci / test (windows-latest) (push) Waiting to run
ci / race (push) Waiting to run
ci / integration (push) Waiting to run
feat: publish initial worker snapshots
2026-07-31 13:22:45 +08:00

178 lines
5.2 KiB
Go

package worker
import (
"context"
"crypto/tls"
"crypto/x509"
"errors"
"fmt"
"net"
"os"
"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
}
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 {
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))
return &Server{listen: controlPlane.Listen, grpcServer: grpcServer, shutdownTimeout: options.ShutdownTimeout}, nil
}
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()
}