package main import ( "context" "fmt" "net" "os" "os/exec" "os/signal" "path/filepath" "syscall" "github.com/spf13/pflag" "golang.org/x/sync/errgroup" "github.com/balabanovds/pgcli/internal/config" "github.com/balabanovds/pgcli/internal/core" "github.com/balabanovds/pgcli/internal/kube" "github.com/balabanovds/pgcli/internal/pg" ) const ( appName = "pgpf" ) func main() { ctx := context.Background() if err := run(ctx); err != nil { fmt.Println(err) os.Exit(1) } } func run(ctx context.Context) error { var defaultAppConfig string if home := core.Homedir(); home != "" { defaultAppConfig = filepath.Join(home, "."+appName, "config.yaml") } var ( appconfig string kubeconfig string namespace string ) pflag.StringVar(&appconfig, "config", defaultAppConfig, "Application configuration file path") pflag.StringVar(&kubeconfig, "kubeconfig", "", "Absolute path to the kubeconfig file") pflag.StringVar(&namespace, "ns", "", "Namespace (required)") pflag.Parse() if namespace == "" { pflag.Usage() os.Exit(1) } appCfg, err := config.InitConfig(appconfig) if err != nil { return fmt.Errorf("failed to initiate application config %s: %w", appconfig, err) } cl, err := kube.NewClient(kubeconfig, &appCfg.Kube) if err != nil { return fmt.Errorf("failed to create Kubernetes client: %w", err) } pods, err := cl.PodNames(ctx, namespace) if err != nil { return fmt.Errorf("failed to get Pods: %w", err) } if len(pods) == 0 { return fmt.Errorf("pods not found") } if len(pods) != 1 { return fmt.Errorf("expected only one pod") } pod := pods[0] secret, err := cl.GetSecret(ctx, pod.Namespace) if err != nil { return fmt.Errorf("failed to get secret for namespace %s: %w", pod.Namespace, err) } // stopCh control the port forwarding lifecycle. When it gets closed the // port forward will terminate stopCh := make(chan struct{}, 1) // readyCh communicate when the port forward is ready to get traffic readyCh := make(chan struct{}) // managing termination signal from the terminal. As you can see the stopCh // gets closed to gracefully handle its termination. sigs := make(chan os.Signal, 1) signal.Notify(sigs, syscall.SIGINT, syscall.SIGTERM) go func() { <-sigs fmt.Println("Bye...") close(stopCh) }() var localport int if appCfg.App.PortForwardOnly.Enabled && appCfg.App.PortForwardOnly.LocalPort != 0 { localport = appCfg.App.PortForwardOnly.LocalPort } else { localport, err = freeTCPPort() if err != nil { return fmt.Errorf("failed to get free local TCP port: %w", err) } } localhost := "127.0.0.1" if appCfg.App.PortForwardOnly.Enabled && appCfg.App.PortForwardOnly.LocalHost != "" { localhost = appCfg.App.PortForwardOnly.LocalHost } var wgerr errgroup.Group wgerr.Go(func() error { return cl.ForwardPort(&kube.ForwardPodRequest{ Pod: pod, LocalPort: localport, PodPort: secret.Port, Stream: core.Stream{ In: os.Stdin, Out: os.Stdout, Err: os.Stderr, }, StopCh: stopCh, ReadyCh: readyCh, }) }) <-readyCh if appCfg.App.PortForwardOnly.Enabled { pgpass, err := pg.NewPgPassHandler() if err != nil { return fmt.Errorf("failed to create pgpass handler: %w", err) } defer func() { if err = pgpass.Close(); err != nil { fmt.Println("error occurred during closing pgpass handler", err) } }() err = pgpass.Append(secret, localhost, localport) if err != nil { return fmt.Errorf("failed to add new data to pgpass file: %w", err) } fmt.Printf("Connection to database %s is ready, temporary data added to .pgpass. Use %s:%d for connection...", secret.DBName, localhost, localport) return wgerr.Wait() } wgerr.Go(func() error { defer close(stopCh) dsn := dsn(secret, localport) var cmd *exec.Cmd switch { case appCfg.App.UseApp.PgCLI: cmd = exec.CommandContext(ctx, "pgcli", dsn) //nolint:gosec case appCfg.App.UseApp.Psql: cmd = exec.CommandContext(ctx, "psql", dsn) //nolint:gosec default: panic("can not run unknown command") } cmd.Stdout = os.Stdout cmd.Stdin = os.Stdin cmd.Stderr = os.Stderr return cmd.Run() }) return wgerr.Wait() } func freeTCPPort() (int, error) { addr, err := net.ResolveTCPAddr("tcp", "localhost:0") if err != nil { return 0, err } l, err := net.ListenTCP("tcp", addr) if err != nil { return 0, err } defer func() { if cerr := l.Close(); cerr != nil { fmt.Printf("failed to close socket: %s", cerr) } }() return l.Addr().(*net.TCPAddr).Port, nil } func dsn(secret *kube.Secret, localport int) string { return fmt.Sprintf("postgresql://%s:%s@localhost:%d/%s", secret.User, secret.Password, localport, secret.DBName) }