Merge pull request #3155 from akhilmhdh/feat/connector

feat: removed pool config from knex and better closing in cli
This commit is contained in:
Maidul Islam
2025-02-27 10:28:26 +09:00
committed by GitHub
3 changed files with 135 additions and 73 deletions
@@ -51,7 +51,6 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
user: providerInputs.username, user: providerInputs.username,
password: providerInputs.password, password: providerInputs.password,
ssl, ssl,
pool: { min: 0, max: 1 },
// @ts-expect-error this is because of knexjs type signature issue. This is directly passed to driver // @ts-expect-error this is because of knexjs type signature issue. This is directly passed to driver
// https://github.com/knex/knex/blob/b6507a7129d2b9fafebf5f831494431e64c6a8a0/lib/dialects/mssql/index.js#L66 // https://github.com/knex/knex/blob/b6507a7129d2b9fafebf5f831494431e64c6a8a0/lib/dialects/mssql/index.js#L66
// https://github.com/tediousjs/tedious/blob/ebb023ed90969a7ec0e4b036533ad52739d921f7/test/config.ci.ts#L19 // https://github.com/tediousjs/tedious/blob/ebb023ed90969a7ec0e4b036533ad52739d921f7/test/config.ci.ts#L19
@@ -106,10 +105,8 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
}; };
if (providerInputs.projectGatewayId) { if (providerInputs.projectGatewayId) {
console.log(">>>>>> inside gateway");
await gatewayProxyWrapper(providerInputs, gatewayCallback); await gatewayProxyWrapper(providerInputs, gatewayCallback);
} else { } else {
console.log(">>>>>> outside gateway");
await gatewayCallback(); await gatewayCallback();
} }
return isConnected; return isConnected;
@@ -121,24 +118,27 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
const password = generatePassword(providerInputs.client); const password = generatePassword(providerInputs.client);
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({ ...providerInputs, port, host });
const { database } = providerInputs; try {
const expiration = new Date(expireAt).toISOString(); const { database } = providerInputs;
const expiration = new Date(expireAt).toISOString();
const creationStatement = handlebars.compile(providerInputs.creationStatement, { noEscape: true })({ const creationStatement = handlebars.compile(providerInputs.creationStatement, { noEscape: true })({
username, username,
password, password,
expiration, expiration,
database database
}); });
const queries = creationStatement.toString().split(";").filter(Boolean); const queries = creationStatement.toString().split(";").filter(Boolean);
await db.transaction(async (tx) => { await db.transaction(async (tx) => {
for (const query of queries) { for (const query of queries) {
// eslint-disable-next-line // eslint-disable-next-line
await tx.raw(query); await tx.raw(query);
} }
}); });
await db.destroy(); } finally {
await db.destroy();
}
}; };
if (providerInputs.projectGatewayId) { if (providerInputs.projectGatewayId) {
await gatewayProxyWrapper(providerInputs, gatewayCallback); await gatewayProxyWrapper(providerInputs, gatewayCallback);
@@ -154,16 +154,18 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
const { database } = providerInputs; const { database } = providerInputs;
const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => { const gatewayCallback = async (host = providerInputs.host, port = providerInputs.port) => {
const db = await $getClient({ ...providerInputs, port, host }); const db = await $getClient({ ...providerInputs, port, host });
const revokeStatement = handlebars.compile(providerInputs.revocationStatement)({ username, database }); try {
const queries = revokeStatement.toString().split(";").filter(Boolean); const revokeStatement = handlebars.compile(providerInputs.revocationStatement)({ username, database });
await db.transaction(async (tx) => { const queries = revokeStatement.toString().split(";").filter(Boolean);
for (const query of queries) { await db.transaction(async (tx) => {
// eslint-disable-next-line for (const query of queries) {
await tx.raw(query); // eslint-disable-next-line
} await tx.raw(query);
}); }
});
await db.destroy(); } finally {
await db.destroy();
}
}; };
if (providerInputs.projectGatewayId) { if (providerInputs.projectGatewayId) {
await gatewayProxyWrapper(providerInputs, gatewayCallback); await gatewayProxyWrapper(providerInputs, gatewayCallback);
@@ -187,18 +189,19 @@ export const SqlDatabaseProvider = ({ gatewayService }: TSqlDatabaseProviderDTO)
expiration, expiration,
database database
}); });
try {
if (renewStatement) { if (renewStatement) {
const queries = renewStatement.toString().split(";").filter(Boolean); const queries = renewStatement.toString().split(";").filter(Boolean);
await db.transaction(async (tx) => { await db.transaction(async (tx) => {
for (const query of queries) { for (const query of queries) {
// eslint-disable-next-line // eslint-disable-next-line
await tx.raw(query); await tx.raw(query);
} }
}); });
}
} finally {
await db.destroy();
} }
await db.destroy();
}; };
if (providerInputs.projectGatewayId) { if (providerInputs.projectGatewayId) {
await gatewayProxyWrapper(providerInputs, gatewayCallback); await gatewayProxyWrapper(providerInputs, gatewayCallback);
+42 -19
View File
@@ -53,34 +53,57 @@ var gatewayCmd = &cobra.Command{
<-sigCh <-sigCh
close(sigStopCh) close(sigStopCh)
cancel() cancel()
// If we get a second signal, force exit
<-sigCh
log.Warn().Msgf("Force exit triggered")
os.Exit(1)
}() }()
// Main gateway retry loop with proper context handling
retryTicker := time.NewTicker(5 * time.Second)
defer retryTicker.Stop()
for { for {
select { if ctx.Err() != nil {
case <-sigStopCh:
log.Info().Msg("Shutting down gateway") log.Info().Msg("Shutting down gateway")
return return
default: }
gatewayInstance, err := gateway.NewGateway(token.Token) gatewayInstance, err := gateway.NewGateway(token.Token)
if err != nil { if err != nil {
util.HandleError(err) util.HandleError(err)
} }
if err = gatewayInstance.ConnectWithRelay(); err != nil { if err = gatewayInstance.ConnectWithRelay(); err != nil {
log.Error().Msgf("Gateway connection error with relay: %s", err) if ctx.Err() != nil {
log.Info().Msg("Restarting gateway...") log.Info().Msg("Shutting down gateway")
time.Sleep(5 * time.Second)
continue
}
err = gatewayInstance.Listen(ctx)
if err == nil {
// meaning everything went smooth and we are exiting
return return
} }
log.Error().Msgf("Gateway listen error: %s", err) log.Error().Msgf("Gateway connection error with relay: %s", err)
log.Info().Msg("Restarting gateway...") log.Info().Msg("Retrying connection in 5 seconds...")
time.Sleep(5 * time.Second) select {
case <-retryTicker.C:
continue
case <-ctx.Done():
log.Info().Msg("Shutting down gateway")
return
}
}
err = gatewayInstance.Listen(ctx)
if ctx.Err() != nil {
log.Info().Msg("Gateway shutdown complete")
return
}
log.Error().Msgf("Gateway listen error: %s", err)
log.Info().Msg("Retrying connection in 5 seconds...")
select {
case <-retryTicker.C:
continue
case <-ctx.Done():
log.Info().Msg("Shutting down gateway")
return
} }
} }
}, },
+50 -14
View File
@@ -112,6 +112,7 @@ func (g *Gateway) Listen(ctx context.Context) error {
} }
log.Info().Msg("Connected with relay") 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.
@@ -139,7 +140,7 @@ func (g *Gateway) Listen(ctx context.Context) error {
g.config.Certificate = gatewayCert.Certificate g.config.Certificate = gatewayCert.Certificate
g.config.CertificateChain = gatewayCert.CertificateChain g.config.CertificateChain = gatewayCert.CertificateChain
done := make(chan bool, 1) shutdownCh := make(chan bool, 1)
if g.config.InfisicalStaticIp != "" { if g.config.InfisicalStaticIp != "" {
log.Info().Msgf("Found static ip from Infisical: %s. Creating permission IP lifecycle", g.config.InfisicalStaticIp) log.Info().Msgf("Found static ip from Infisical: %s. Creating permission IP lifecycle", g.config.InfisicalStaticIp)
@@ -150,7 +151,7 @@ func (g *Gateway) Listen(ctx context.Context) error {
g.registerPermissionLifecycle(func() error { g.registerPermissionLifecycle(func() error {
err := relayNonTlsConn.CreatePermissions(peerAddr) err := relayNonTlsConn.CreatePermissions(peerAddr)
return err return err
}, done) }, shutdownCh)
} }
cert, err := tls.X509KeyPair([]byte(gatewayCert.Certificate), []byte(gatewayCert.PrivateKey)) cert, err := tls.X509KeyPair([]byte(gatewayCert.Certificate), []byte(gatewayCert.PrivateKey))
@@ -170,21 +171,32 @@ func (g *Gateway) Listen(ctx context.Context) error {
errCh := make(chan error, 1) errCh := make(chan error, 1)
log.Info().Msg("Gateway started successfully") log.Info().Msg("Gateway started successfully")
g.registerHeartBeat(errCh, done) g.registerHeartBeat(errCh, shutdownCh)
g.registerRelayIsActive(relayNonTlsConn.Addr().String(), errCh, done) g.registerRelayIsActive(relayNonTlsConn.Addr().String(), errCh, shutdownCh)
// Create a WaitGroup to track active connections // Create a WaitGroup to track active connections
var wg sync.WaitGroup var wg sync.WaitGroup
go func() { go func() {
for { for {
if relayDeadlineConn, ok := relayConn.(*net.TCPListener); ok {
relayDeadlineConn.SetDeadline(time.Now().Add(1 * time.Second))
}
select { select {
case <-done: case <-ctx.Done():
return
case <-shutdownCh:
return return
default: default:
// Accept new relay connection // Accept new relay connection
conn, err := relayConn.Accept() conn, err := relayConn.Accept()
if err != nil { if err != nil {
// Check if it's a timeout error (which we expect due to our deadline)
if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
continue
}
if !strings.Contains(err.Error(), "data contains incomplete STUN or TURN frame") { if !strings.Contains(err.Error(), "data contains incomplete STUN or TURN frame") {
log.Error().Msgf("Failed to accept connection: %v", err) log.Error().Msgf("Failed to accept connection: %v", err)
} }
@@ -198,7 +210,11 @@ func (g *Gateway) Listen(ctx context.Context) error {
continue continue
} }
// Set a deadline for the handshake to prevent hanging
tlsConn.SetDeadline(time.Now().Add(10 * time.Second))
err = tlsConn.Handshake() err = tlsConn.Handshake()
// Clear the deadline after handshake
tlsConn.SetDeadline(time.Time{})
if err != nil { if err != nil {
log.Error().Msgf("TLS handshake failed: %v", err) log.Error().Msgf("TLS handshake failed: %v", err)
conn.Close() conn.Close()
@@ -212,34 +228,54 @@ func (g *Gateway) Listen(ctx context.Context) error {
commonName := state.PeerCertificates[0].Subject.CommonName commonName := state.PeerCertificates[0].Subject.CommonName
if organizationUnit[0] != "gateway-client" || commonName != "cloud" { if organizationUnit[0] != "gateway-client" || commonName != "cloud" {
log.Error().Msgf("Client certificate verification failed. Received %s, %s", organizationUnit, commonName) log.Error().Msgf("Client certificate verification failed. Received %s, %s", organizationUnit, commonName)
conn.Close()
continue continue
} }
} }
// Handle the connection in a goroutine // Handle the connection in a goroutine
wg.Add(1) wg.Add(1)
go func() { go func(c net.Conn) {
defer wg.Done() defer wg.Done()
handleConnection(conn) defer c.Close()
}()
// Monitor parent context to close this connection when needed
go func() {
select {
case <-ctx.Done():
c.Close() // Force close connection when context is canceled
case <-shutdownCh:
c.Close() // Force close connection when accepting loop is done
}
}()
handleConnection(c)
}(conn)
} }
} }
}() }()
var isShutdown bool
select { select {
case <-ctx.Done(): case <-ctx.Done():
log.Info().Msg("Shutting down gateway...") log.Info().Msg("Shutting down gateway...")
isShutdown = true
case err = <-errCh: case err = <-errCh:
} }
// Signal the accept loop to stop // Signal the accept loop to stop
close(done) close(shutdownCh)
wg.Wait()
if isShutdown { // Set a timeout for waiting on connections to close
log.Info().Msg("Gateway shutdown complete") waitCh := make(chan struct{})
go func() {
wg.Wait()
close(waitCh)
}()
select {
case <-waitCh:
// All connections closed normally
case <-time.After(5 * time.Second):
log.Warn().Msg("Timeout waiting for connections to close gracefully")
} }
return err return err