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" "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) } 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, GetCertificate: certificate.ServerCertificate, 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() }