pgcli/cmd/pgpf/main.go

200 lines
4.6 KiB
Go

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)
}