feat: review feedback, and more changes in cli

This commit is contained in:
=
2025-02-20 23:37:19 +05:30
parent f34370cb9d
commit f7406ea8f8
16 changed files with 336 additions and 82 deletions

View File

@@ -583,3 +583,20 @@ func CallExchangeRelayCertV1(httpClient *resty.Client, request ExchangeRelayCert
return &resBody, nil
}
func CallGatewayHeartBeatV1(httpClient *resty.Client) error {
response, err := httpClient.
R().
SetHeader("User-Agent", USER_AGENT).
Post(fmt.Sprintf("%v/v1/gateways/heartbeat", config.INFISICAL_URL))
if err != nil {
return fmt.Errorf("CallGatewayHeartBeatV1: Unable to complete api request [err=%w]", err)
}
if response.IsError() {
return fmt.Errorf("CallGatewayHeartBeatV1: Unsuccessful response [%v %v] [status-code=%v] [response=%v]", response.Request.Method, response.Request.URL, response.StatusCode(), response.String())
}
return nil
}

View File

@@ -5,10 +5,17 @@ import (
// "github.com/Infisical/infisical-merge/packages/api"
// "github.com/Infisical/infisical-merge/packages/models"
"context"
"fmt"
"os"
"os/signal"
"syscall"
"time"
"github.com/Infisical/infisical-merge/packages/gateway"
"github.com/Infisical/infisical-merge/packages/util"
"github.com/rs/zerolog/log"
// "github.com/Infisical/infisical-merge/packages/visualize"
// "github.com/rs/zerolog/log"
@@ -33,20 +40,48 @@ var gatewayCmd = &cobra.Command{
util.HandleError(fmt.Errorf("Token not found"))
}
gatewayInstance, err := gateway.NewGateway(token.Token)
if err != nil {
util.HandleError(err)
}
if err = gatewayInstance.ConnectWithRelay(); err != nil {
util.HandleError(err)
}
if err := gatewayInstance.Listen(); err != nil {
util.HandleError(err)
}
Telemetry.CaptureEvent("cli-command:gateway", posthog.NewProperties().Set("version", util.CLI_VERSION))
sigCh := make(chan os.Signal, 1)
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
sigStopCh := make(chan bool, 1)
go func() {
<-sigCh
close(sigStopCh)
}()
ctx, cancel := context.WithCancel(cmd.Context())
defer cancel()
for {
select {
case <-sigStopCh:
log.Info().Msg("Shutting down gateway")
return
default:
gatewayInstance, err := gateway.NewGateway(token.Token)
if err != nil {
util.HandleError(err)
}
if err = gatewayInstance.ConnectWithRelay(); err != nil {
log.Error().Msgf("Gateway connection error with relay: %s", err)
log.Info().Msg("Restarting gateway...")
time.Sleep(5 * time.Second)
continue
}
if err := gatewayInstance.Listen(ctx); err == nil {
// meaning everything went smooth and we are exiting
return
}
log.Error().Msgf("Gateway listen error: %s", err)
log.Info().Msg("Restarting gateway...")
time.Sleep(5 * time.Second)
}
}
},
}

View File

@@ -59,7 +59,9 @@ func handleConnection(conn net.Conn) {
CopyData(conn, destTarget)
return
case "PING":
conn.Write([]byte("PONG"))
if _, err := conn.Write([]byte("PONG")); err != nil {
log.Error().Msgf("Error writing PONG response: %v", err)
}
return
default:
log.Error().Msgf("Unknown command: %s", string(cmd))

View File

@@ -1,11 +1,13 @@
package gateway
import (
"context"
"crypto/tls"
"crypto/x509"
"fmt"
"net"
"strings"
"sync"
"time"
"github.com/Infisical/infisical-merge/packages/api"
@@ -91,12 +93,13 @@ func (g *Gateway) ConnectWithRelay() error {
return nil
}
func (g *Gateway) Listen() error {
func (g *Gateway) Listen(ctx context.Context) error {
defer g.client.Close()
err := g.client.Listen()
if err != nil {
return fmt.Errorf("Failed to listen to relay server: %w", err)
}
log.Info().Msg("Connected with relay")
// Allocate a relay socket on the TURN server. On success, it
// will return a net.PacketConn which represents the remote
@@ -106,6 +109,7 @@ func (g *Gateway) Listen() error {
return fmt.Errorf("Failed to allocate relay connection: %w", err)
}
log.Info().Msg(relayNonTlsConn.Addr().String())
defer func() {
if closeErr := relayNonTlsConn.Close(); closeErr != nil {
log.Error().Msgf("Failed to close connection: %s", closeErr)
@@ -128,20 +132,12 @@ func (g *Gateway) Listen() error {
g.config.Certificate = gatewayCert.Certificate
g.config.CertificateChain = gatewayCert.CertificateChain
go func() {
done := make(chan bool, 1)
g.registerPermissionLifecycle(func() error {
err := relayNonTlsConn.CreatePermissions(peerAddr)
if err != nil {
log.Error().Msgf("Failed to refresh permission: %s", err)
}
log.Printf("Created permission for incoming connections")
ticker := time.NewTicker(2 * time.Minute) // Refresh before 5-min expiry
for range ticker.C {
err := relayNonTlsConn.CreatePermissions(peerAddr)
if err != nil {
log.Error().Msgf("Failed to refresh permission: %s", err)
}
}
}()
return err
}, done)
cert, err := tls.X509KeyPair([]byte(gatewayCert.Certificate), []byte(gatewayCert.PrivateKey))
if err != nil {
@@ -158,41 +154,146 @@ func (g *Gateway) Listen() error {
ClientAuth: tls.RequireAndVerifyClientCert,
})
errCh := make(chan error, 1)
log.Info().Msg("Connector started successfully")
for {
// Accept new relay connection
conn, err := relayConn.Accept()
if err != nil {
log.Error().Msgf("Failed to accept connection: %v", err)
continue
}
g.registerHeartBeat(errCh, done)
g.registerRelayIsActive(relayNonTlsConn.Addr().String(), errCh, done)
tlsConn, ok := conn.(*tls.Conn)
if !ok {
log.Error().Msg("Failed to convert to TLS connection")
conn.Close()
continue
}
// Create a WaitGroup to track active connections
var wg sync.WaitGroup
err = tlsConn.Handshake()
if err != nil {
log.Error().Msgf("TLS handshake failed: %v", err)
conn.Close()
continue
}
go func() {
for {
select {
case <-done:
return
default:
// Accept new relay connection
conn, err := relayConn.Accept()
if err != nil {
if !strings.Contains(err.Error(), "data contains incomplete STUN or TURN frame") {
log.Error().Msgf("Failed to accept connection: %v", err)
}
continue
}
// Get connection state which contains certificate information
state := tlsConn.ConnectionState()
if len(state.PeerCertificates) > 0 {
organizationUnit := state.PeerCertificates[0].Subject.OrganizationalUnit
commonName := state.PeerCertificates[0].Subject.CommonName
if organizationUnit[0] != "gateway-client" && commonName != "cloud" {
log.Error().Msgf("Client certificate verification failed. Received %s, %s", organizationUnit, commonName)
continue
tlsConn, ok := conn.(*tls.Conn)
if !ok {
log.Error().Msg("Failed to convert to TLS connection")
conn.Close()
continue
}
err = tlsConn.Handshake()
if err != nil {
log.Error().Msgf("TLS handshake failed: %v", err)
conn.Close()
continue
}
// Get connection state which contains certificate information
state := tlsConn.ConnectionState()
if len(state.PeerCertificates) > 0 {
organizationUnit := state.PeerCertificates[0].Subject.OrganizationalUnit
commonName := state.PeerCertificates[0].Subject.CommonName
if organizationUnit[0] != "gateway-client" && commonName != "cloud" {
log.Error().Msgf("Client certificate verification failed. Received %s, %s", organizationUnit, commonName)
continue
}
}
// Handle the connection in a goroutine
wg.Add(1)
go func() {
defer wg.Done()
handleConnection(conn)
}()
}
}
}()
// Handle the connection in a goroutine
go handleConnection(conn)
var isShutdown bool
select {
case <-ctx.Done():
log.Info().Msg("Shutting down gateway...")
isShutdown = true
case err = <-errCh:
}
// Signal the accept loop to stop
close(done)
wg.Wait()
if isShutdown {
log.Info().Msg("Gateway shutdown complete")
}
return err
}
func (g *Gateway) registerHeartBeat(errCh chan error, done chan bool) {
ticker := time.NewTicker(1 * time.Hour)
go func() {
// wait for 5 mins
time.Sleep(5 * time.Second)
err := api.CallGatewayHeartBeatV1(g.httpClient)
if err != nil {
log.Error().Msgf("Failed to register heartbeat: %s", err)
}
for {
select {
case <-done:
ticker.Stop()
return
case <-ticker.C:
err := api.CallGatewayHeartBeatV1(g.httpClient)
errCh <- err
}
}
}()
}
func (g *Gateway) registerPermissionLifecycle(permissionFn func() error, done chan bool) {
ticker := time.NewTicker(3 * time.Minute)
go func() {
// wait for 5 mins
permissionFn()
log.Printf("Ceated permission for incoming connections")
for {
select {
case <-done:
ticker.Stop()
return
case <-ticker.C:
permissionFn()
}
}
}()
}
func (g *Gateway) registerRelayIsActive(serverAddr string, errCh chan error, done chan bool) {
ticker := time.NewTicker(10 * time.Second)
go func() {
time.Sleep(5 * time.Second)
for {
select {
case <-done:
ticker.Stop()
return
case <-ticker.C:
conn, err := net.Dial("tcp", serverAddr)
if err != nil {
errCh <- err
return
}
if conn != nil {
conn.Close()
}
}
}
}()
}