Files
query-orchestration/internal/database/connection.go
T

136 lines
2.6 KiB
Go
Raw Normal View History

2024-12-18 18:54:48 +00:00
package database
import (
"context"
"fmt"
"log"
2025-01-14 17:28:26 +00:00
"net/url"
2025-01-23 14:56:20 +00:00
"queryorchestration/internal/database/repository"
"queryorchestration/internal/server/env"
2024-12-18 18:54:48 +00:00
2025-01-14 17:28:26 +00:00
"github.com/docker/go-connections/nat"
2024-12-18 18:54:48 +00:00
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
type config struct {
user string
password string
host string
2025-01-14 17:28:26 +00:00
port nat.Port
name string
disableSSL bool
}
func mustGetConfig() *config {
user := env.GetPanic("DB_USER")
pass := env.GetPanic("DB_PASS")
host := env.GetPanic("DB_HOST")
port := env.GetPanic("DB_PORT")
2025-01-14 17:28:26 +00:00
natPort, err := nat.NewPort("tcp", port)
if err != nil {
log.Panicf("Failed to create port: %v", err)
}
name := env.GetPanic("DB_NAME")
disableSSL, _ := env.Get("DB_NOSSL")
return &config{
user: user,
password: pass,
host: host,
2025-01-14 17:28:26 +00:00
port: natPort,
name: name,
disableSSL: disableSSL == "1",
}
}
2025-01-14 17:28:26 +00:00
func mustGetURI() *url.URL {
2024-12-18 18:54:48 +00:00
driver := "postgres"
conf := mustGetConfig()
opts := ""
if conf.disableSSL {
opts += "sslmode=disable"
}
2025-01-14 17:28:26 +00:00
connStr := fmt.Sprintf("%s://%s:%s@%s:%d/%s?%s", driver, conf.user, conf.password, conf.host, conf.port.Int(), conf.name, opts)
uri, err := url.Parse(connStr)
if err != nil {
log.Panicf("Unable to parse URI: %s", err)
}
return uri
}
func getPoolConfig() (*pgxpool.Config, error) {
2025-01-14 17:28:26 +00:00
connStr := mustGetURI()
2025-01-14 17:28:26 +00:00
config, err := pgxpool.ParseConfig(connStr.String())
if err != nil {
return nil, err
}
enumTypes := []string{
"querytype",
}
config.AfterConnect = func(ctx context.Context, conn *pgx.Conn) error {
for _, enumName := range enumTypes {
dt, err := conn.LoadType(ctx, enumName)
if err != nil {
return fmt.Errorf("failed to load suitable enum type: %w", err)
}
conn.TypeMap().RegisterType(dt)
}
return nil
}
return config, nil
}
func mustGetPoolConfig() *pgxpool.Config {
config, err := getPoolConfig()
if err != nil {
log.Panicf("Unable to create database config: %v\n", err)
}
return config
2024-12-18 18:54:48 +00:00
}
func GetDBPool(ctx context.Context) *pgxpool.Pool {
config := mustGetPoolConfig()
2024-12-18 18:54:48 +00:00
pool, err := pgxpool.NewWithConfig(ctx, config)
2024-12-18 18:54:48 +00:00
if err != nil {
2024-12-24 18:11:25 +00:00
log.Panicf("Unable to create database pool: %v\n", err)
2024-12-18 18:54:48 +00:00
}
return pool
}
2025-01-23 14:56:20 +00:00
func ExecuteTransaction(ctx context.Context, db *Connection, executeQueries func(context.Context, *repository.Queries) error) error {
tx, err := db.Pool.Begin(ctx)
if err != nil {
return err
}
defer func() {
_ = tx.Rollback(ctx)
}()
qtx := db.Queries.WithTx(tx)
err = executeQueries(ctx, qtx)
if err != nil {
return err
}
err = tx.Commit(ctx)
if err != nil {
return err
}
return nil
}