package database import ( "context" "fmt" "log/slog" "time" "queryorchestration/internal/database/repository" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) type Pool interface { repository.DBTX Begin(ctx context.Context) (pgx.Tx, error) Ping(ctx context.Context) error } func (b *DBConfig) GetDBPool() Pool { return b.DBPool } func (b *DBConfig) GetDBQueries() *repository.Queries { return b.DBQueries } func (b *DBConfig) SetDBPoolConfig() error { config, err := pgxpool.ParseConfig(b.GetDBURI()) if err != nil { return err } b.DBEnumTypes = []string{ "querytype", } config.AfterConnect = func(ctx context.Context, conn *pgx.Conn) error { for _, enumName := range b.DBEnumTypes { 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 } b.DBPoolConfig = config return nil } func (b *DBConfig) SetDBPool(ctx context.Context) error { err := b.SetDBPoolConfig() if err != nil { return err } pool, err := pgxpool.NewWithConfig(ctx, b.DBPoolConfig) if err != nil { return err } b.DBPool = pool b.DBQueries = repository.New(b.DBPool) err = b.DBPing(ctx) if err != nil { return err } return nil } func (b *DBConfig) DBPing(ctx context.Context) error { timeout := time.After(30 * time.Second) tick := time.NewTicker(500 * time.Millisecond) defer tick.Stop() for { select { case <-timeout: return fmt.Errorf("unable to ping database") case <-tick.C: err := b.DBPool.Ping(ctx) if err == nil { slog.Info("Successful database ping") return nil } slog.Debug("Attempted database ping", "host", b.DBHost, "port", b.DBPort) case <-ctx.Done(): return ctx.Err() } } }