mirror of
https://github.com/awatertrevi/infisical.git
synced 2026-10-07 14:27:30 +00:00
feat: completed cli for gateway
This commit is contained in:
@@ -3,7 +3,7 @@ package gateway
|
|||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
"bytes"
|
"bytes"
|
||||||
"fmt"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -20,6 +20,9 @@ func handleConnection(conn net.Conn) {
|
|||||||
for {
|
for {
|
||||||
msg, err := reader.ReadBytes('\n')
|
msg, err := reader.ReadBytes('\n')
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if errors.Is(err, io.EOF) {
|
||||||
|
return
|
||||||
|
}
|
||||||
log.Error().Msgf("Error reading command: %s", err)
|
log.Error().Msgf("Error reading command: %s", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -30,9 +33,7 @@ func handleConnection(conn net.Conn) {
|
|||||||
switch string(cmd) {
|
switch string(cmd) {
|
||||||
case "FORWARD-TCP":
|
case "FORWARD-TCP":
|
||||||
proxyAddress := string(bytes.Split(args, []byte(" "))[0])
|
proxyAddress := string(bytes.Split(args, []byte(" "))[0])
|
||||||
fmt.Println(proxyAddress)
|
|
||||||
destTarget, err := net.Dial("tcp", proxyAddress)
|
destTarget, err := net.Dial("tcp", proxyAddress)
|
||||||
fmt.Println(err)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
log.Error().Msgf("Failed to connect to target: %v", err)
|
log.Error().Msgf("Failed to connect to target: %v", err)
|
||||||
return
|
return
|
||||||
@@ -56,12 +57,13 @@ func handleConnection(conn net.Conn) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
CopyData(conn, destTarget)
|
CopyData(conn, destTarget)
|
||||||
break
|
return
|
||||||
case "PING":
|
case "PING":
|
||||||
conn.Write([]byte("PONG\n"))
|
conn.Write([]byte("PONG"))
|
||||||
|
return
|
||||||
default:
|
default:
|
||||||
log.Error().Msgf("Unknown command: %s", string(cmd))
|
log.Error().Msgf("Unknown command: %s", string(cmd))
|
||||||
break
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -71,39 +73,33 @@ type CloseWrite interface {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func CopyData(src, dst net.Conn) {
|
func CopyData(src, dst net.Conn) {
|
||||||
// Create a WaitGroup to wait for both copy operations
|
|
||||||
var wg sync.WaitGroup
|
var wg sync.WaitGroup
|
||||||
wg.Add(2)
|
wg.Add(2)
|
||||||
|
|
||||||
// Start copying in both directions
|
copyAndClose := func(dst, src net.Conn, done chan<- bool) {
|
||||||
go func() {
|
|
||||||
defer wg.Done()
|
defer wg.Done()
|
||||||
if _, err := io.Copy(dst, src); err != nil {
|
_, err := io.Copy(dst, src)
|
||||||
log.Error().Msgf("Error copying postgres->client: %v", err)
|
if err != nil && !errors.Is(err, io.EOF) {
|
||||||
|
log.Error().Msgf("Copy error: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if e, ok := dst.(CloseWrite); ok {
|
// Signal we're done writing
|
||||||
log.Print("Closing dst")
|
done <- true
|
||||||
e.CloseWrite()
|
|
||||||
} else {
|
|
||||||
|
|
||||||
log.Print("Not closed")
|
// Half close the connection if possible
|
||||||
|
if c, ok := dst.(CloseWrite); ok {
|
||||||
|
c.CloseWrite()
|
||||||
}
|
}
|
||||||
}()
|
}
|
||||||
|
|
||||||
go func() {
|
done1 := make(chan bool, 1)
|
||||||
defer wg.Done()
|
done2 := make(chan bool, 1)
|
||||||
if _, err := io.Copy(src, dst); err != nil {
|
|
||||||
log.Error().Msgf("Error copying client->postgres: %v", err)
|
go copyAndClose(dst, src, done1)
|
||||||
}
|
go copyAndClose(src, dst, done2)
|
||||||
if e, ok := src.(CloseWrite); ok {
|
|
||||||
log.Print("Closing src")
|
|
||||||
e.CloseWrite()
|
|
||||||
} else {
|
|
||||||
log.Print("Not closed")
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Wait for both copies to complete
|
// Wait for both copies to complete
|
||||||
|
<-done1
|
||||||
|
<-done2
|
||||||
wg.Wait()
|
wg.Wait()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -48,16 +48,15 @@ func (g *Gateway) ConnectWithRelay() error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Dial TURN Server
|
turnServerAddr, err := net.ResolveTCPAddr("tcp", relayDetails.TurnServerAddress)
|
||||||
conn, err := net.Dial("tcp", relayDetails.TurnServerAddress)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("Failed to connect with relay server: %w", err)
|
return fmt.Errorf("Failed to resolve TURN server address: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if tcpConn, ok := conn.(*net.TCPConn); ok {
|
// Dial TURN Server
|
||||||
tcpConn.SetKeepAlive(true)
|
conn, err := net.DialTCP("tcp", nil, turnServerAddr)
|
||||||
tcpConn.SetKeepAlivePeriod(10 * time.Second)
|
if err != nil {
|
||||||
tcpConn.SetNoDelay(true)
|
return fmt.Errorf("Failed to connect with relay server: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Start a new TURN Client and wrap our net.Conn in a STUNConn
|
// Start a new TURN Client and wrap our net.Conn in a STUNConn
|
||||||
@@ -77,11 +76,6 @@ func (g *Gateway) ConnectWithRelay() error {
|
|||||||
return fmt.Errorf("Failed to create relay client: %w", err)
|
return fmt.Errorf("Failed to create relay client: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = client.Listen()
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("Failed to listen to relay server: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
g.config = &GatewayConfig{
|
g.config = &GatewayConfig{
|
||||||
TurnServerUsername: relayDetails.TurnServerUsername,
|
TurnServerUsername: relayDetails.TurnServerUsername,
|
||||||
TurnServerPassword: relayDetails.TurnServerPassword,
|
TurnServerPassword: relayDetails.TurnServerPassword,
|
||||||
@@ -99,6 +93,11 @@ func (g *Gateway) ConnectWithRelay() error {
|
|||||||
|
|
||||||
func (g *Gateway) Listen() error {
|
func (g *Gateway) Listen() error {
|
||||||
defer g.client.Close()
|
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
|
// Allocate a relay socket on the TURN server. On success, it
|
||||||
// will return a net.PacketConn which represents the remote
|
// will return a net.PacketConn which represents the remote
|
||||||
// socket.
|
// socket.
|
||||||
@@ -106,6 +105,7 @@ func (g *Gateway) Listen() error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("Failed to allocate relay connection: %w", err)
|
return fmt.Errorf("Failed to allocate relay connection: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
if closeErr := relayNonTlsConn.Close(); closeErr != nil {
|
if closeErr := relayNonTlsConn.Close(); closeErr != nil {
|
||||||
log.Error().Msgf("Failed to close connection: %s", closeErr)
|
log.Error().Msgf("Failed to close connection: %s", closeErr)
|
||||||
@@ -148,7 +148,6 @@ func (g *Gateway) Listen() error {
|
|||||||
return fmt.Errorf("failed to parse cert: %s", err)
|
return fmt.Errorf("failed to parse cert: %s", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Println(relayNonTlsConn.Addr().String())
|
|
||||||
caCertPool := x509.NewCertPool()
|
caCertPool := x509.NewCertPool()
|
||||||
caCertPool.AppendCertsFromPEM([]byte(gatewayCert.CertificateChain))
|
caCertPool.AppendCertsFromPEM([]byte(gatewayCert.CertificateChain))
|
||||||
|
|
||||||
@@ -159,8 +158,8 @@ func (g *Gateway) Listen() error {
|
|||||||
ClientAuth: tls.RequireAndVerifyClientCert,
|
ClientAuth: tls.RequireAndVerifyClientCert,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
log.Info().Msg("Connector started successfully")
|
||||||
for {
|
for {
|
||||||
log.Info().Msg("Connector started successfully")
|
|
||||||
// Accept new relay connection
|
// Accept new relay connection
|
||||||
conn, err := relayConn.Accept()
|
conn, err := relayConn.Accept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -168,6 +167,31 @@ func (g *Gateway) Listen() error {
|
|||||||
continue
|
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
|
// Handle the connection in a goroutine
|
||||||
go handleConnection(conn)
|
go handleConnection(conn)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user