mirror of
https://github.com/awatertrevi/infisical.git
synced 2026-10-03 23:26:09 +00:00
Merge pull request #3155 from akhilmhdh/feat/connector
feat: removed pool config from knex and better closing in cli
This commit is contained in:
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user