fix(go.sum): update ResolveSpec dependency to v1.0.87
CI / build-and-test (push) Failing after 1s
Release / release (push) Failing after 19m26s

This commit is contained in:
Hein
2026-06-23 13:17:16 +02:00
parent 0227912325
commit 1adf50e3db
2436 changed files with 1078758 additions and 114 deletions
+24
View File
@@ -0,0 +1,24 @@
Copyright (c) 2021 Vladimir Mihailenco. All rights reserved.
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are
met:
* Redistributions of source code must retain the above copyright
notice, this list of conditions and the following disclaimer.
* Redistributions in binary form must reproduce the above
copyright notice, this list of conditions and the following disclaimer
in the documentation and/or other materials provided with the
distribution.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR
A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT
OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL,
SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT
LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY
THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+38
View File
@@ -0,0 +1,38 @@
# pgdriver
[![PkgGoDev](https://pkg.go.dev/badge/github.com/uptrace/bun/driver/pgdriver)](https://pkg.go.dev/github.com/uptrace/bun/driver/pgdriver)
pgdriver is a database/sql driver for PostgreSQL based on [go-pg](https://github.com/go-pg/pg) code.
You can install it with:
```shell
go get github.com/uptrace/bun/driver/pgdriver
```
And then create a `sql.DB` using it:
```go
import _ "github.com/uptrace/bun/driver/pgdriver"
dsn := "postgres://postgres:@localhost:5432/test"
db, err := sql.Open("pg", dsn)
```
Alternatively:
```go
dsn := "postgres://postgres:@localhost:5432/test"
db := sql.OpenDB(pgdriver.NewConnector(pgdriver.WithDSN(dsn)))
```
[Benchmark](https://github.com/go-bun/bun-benchmark):
```
BenchmarkInsert/pg-12 7254 148380 ns/op 900 B/op 13 allocs/op
BenchmarkInsert/pgx-12 6494 166391 ns/op 2076 B/op 26 allocs/op
BenchmarkSelect/pg-12 9100 132952 ns/op 1417 B/op 18 allocs/op
BenchmarkSelect/pgx-12 8199 154920 ns/op 3679 B/op 60 allocs/op
```
See [documentation](https://bun.uptrace.dev/postgres/) for more details.
+193
View File
@@ -0,0 +1,193 @@
package pgdriver
import (
"encoding/hex"
"fmt"
"io"
"strconv"
"strings"
"time"
)
const (
pgBool = 16
pgInt2 = 21
pgInt4 = 23
pgInt8 = 20
pgFloat4 = 700
pgFloat8 = 701
pgText = 25
pgVarchar = 1043
pgBytea = 17
pgDate = 1082
pgTimestamp = 1114
pgTimestamptz = 1184
)
func readColumnValue(rd *reader, dataType int32, dataLen int) (interface{}, error) {
if dataLen == -1 {
return nil, nil
}
switch dataType {
case pgBool:
return readBoolCol(rd, dataLen)
case pgInt2:
return readIntCol(rd, dataLen, 16)
case pgInt4:
return readIntCol(rd, dataLen, 32)
case pgInt8:
return readIntCol(rd, dataLen, 64)
case pgFloat4:
return readFloatCol(rd, dataLen, 32)
case pgFloat8:
return readFloatCol(rd, dataLen, 64)
case pgTimestamp:
return readTimeCol(rd, dataLen)
case pgTimestamptz:
return readTimeCol(rd, dataLen)
case pgDate:
// Return a string and let the scanner to convert string to time.Time if necessary.
return readStringCol(rd, dataLen)
case pgText, pgVarchar:
return readStringCol(rd, dataLen)
case pgBytea:
return readBytesCol(rd, dataLen)
}
b := make([]byte, dataLen)
if _, err := io.ReadFull(rd, b); err != nil {
return nil, err
}
return b, nil
}
func readBoolCol(rd *reader, n int) (interface{}, error) {
tmp, err := rd.ReadTemp(n)
if err != nil {
return nil, err
}
return len(tmp) == 1 && (tmp[0] == 't' || tmp[0] == '1'), nil
}
func readIntCol(rd *reader, n int, bitSize int) (interface{}, error) {
if n <= 0 {
return 0, nil
}
tmp, err := rd.ReadTemp(n)
if err != nil {
return 0, err
}
return strconv.ParseInt(bytesToString(tmp), 10, bitSize)
}
func readFloatCol(rd *reader, n int, bitSize int) (interface{}, error) {
if n <= 0 {
return 0, nil
}
tmp, err := rd.ReadTemp(n)
if err != nil {
return 0, err
}
return strconv.ParseFloat(bytesToString(tmp), bitSize)
}
func readStringCol(rd *reader, n int) (interface{}, error) {
if n <= 0 {
return "", nil
}
b := make([]byte, n)
if _, err := io.ReadFull(rd, b); err != nil {
return nil, err
}
return bytesToString(b), nil
}
func readBytesCol(rd *reader, n int) (interface{}, error) {
if n <= 0 {
return []byte{}, nil
}
tmp, err := rd.ReadTemp(n)
if err != nil {
return nil, err
}
if len(tmp) < 2 || tmp[0] != '\\' || tmp[1] != 'x' {
return nil, fmt.Errorf("pgdriver: can't parse bytea: %q", tmp)
}
tmp = tmp[2:] // Cut off "\x".
b := make([]byte, hex.DecodedLen(len(tmp)))
if _, err := hex.Decode(b, tmp); err != nil {
return nil, err
}
return b, nil
}
func readTimeCol(rd *reader, n int) (interface{}, error) {
if n <= 0 {
return time.Time{}, nil
}
tmp, err := rd.ReadTemp(n)
if err != nil {
return time.Time{}, err
}
tm, err := ParseTime(bytesToString(tmp))
if err != nil {
return time.Time{}, err
}
return tm, nil
}
const (
dateFormat = "2006-01-02"
timeFormat = "15:04:05.999999999"
timestampFormat = "2006-01-02 15:04:05.999999999"
timestamptzFormat = "2006-01-02 15:04:05.999999999-07:00:00"
timestamptzFormat2 = "2006-01-02 15:04:05.999999999-07:00"
timestamptzFormat3 = "2006-01-02 15:04:05.999999999-07"
)
func ParseTime(s string) (time.Time, error) {
switch l := len(s); {
case l < len("15:04:05"):
return time.Time{}, fmt.Errorf("pgdriver: can't parse time=%q", s)
case l <= len(timeFormat):
if s[2] == ':' {
return time.ParseInLocation(timeFormat, s, time.UTC)
}
return time.ParseInLocation(dateFormat, s, time.UTC)
default:
if s[10] == 'T' {
return time.Parse(time.RFC3339Nano, s)
}
if c := s[l-9]; c == '+' || c == '-' {
return time.Parse(timestamptzFormat, s)
}
if c := s[l-6]; c == '+' || c == '-' {
return time.Parse(timestamptzFormat2, s)
}
if c := s[l-3]; c == '+' || c == '-' {
if strings.HasSuffix(s, "+00") {
s = s[:len(s)-3]
return time.ParseInLocation(timestampFormat, s, time.UTC)
}
return time.Parse(timestamptzFormat3, s)
}
return time.ParseInLocation(timestampFormat, s, time.UTC)
}
}
+419
View File
@@ -0,0 +1,419 @@
package pgdriver
import (
"context"
"crypto/tls"
"crypto/x509"
"errors"
"fmt"
"io/ioutil"
"net"
"net/url"
"os"
"strconv"
"strings"
"time"
)
type Config struct {
// Network type, either tcp or unix.
// Default is tcp.
Network string
// TCP host:port or Unix socket depending on Network.
Addr string
// Dial timeout for establishing new connections.
// Default is 5 seconds.
DialTimeout time.Duration
// Dialer creates new network connection and has priority over
// Network and Addr options.
Dialer func(ctx context.Context, network, addr string) (net.Conn, error)
// TLS config for secure connections.
TLSConfig *tls.Config
User string
Password string
Database string
AppName string
// PostgreSQL session parameters updated with `SET` command when a connection is created.
ConnParams map[string]interface{}
// Timeout for socket reads. If reached, commands fail with a timeout instead of blocking.
ReadTimeout time.Duration
// Timeout for socket writes. If reached, commands fail with a timeout instead of blocking.
WriteTimeout time.Duration
// ResetSessionFunc is called prior to executing a query on a connection that has been used before.
ResetSessionFunc func(context.Context, *Conn) error
}
func newDefaultConfig() *Config {
host := env("PGHOST", "localhost")
port := env("PGPORT", "5432")
cfg := &Config{
Network: "tcp",
Addr: net.JoinHostPort(host, port),
DialTimeout: 5 * time.Second,
TLSConfig: &tls.Config{InsecureSkipVerify: true},
User: env("PGUSER", "postgres"),
Database: env("PGDATABASE", "postgres"),
ReadTimeout: 10 * time.Second,
WriteTimeout: 5 * time.Second,
}
cfg.Dialer = func(ctx context.Context, network, addr string) (net.Conn, error) {
netDialer := &net.Dialer{
Timeout: cfg.DialTimeout,
KeepAlive: 5 * time.Minute,
}
return netDialer.DialContext(ctx, network, addr)
}
return cfg
}
type Option func(cfg *Config)
// Deprecated. Use Option instead.
type DriverOption = Option
func WithNetwork(network string) Option {
if network == "" {
panic("network is empty")
}
return func(cfg *Config) {
cfg.Network = network
}
}
func WithAddr(addr string) Option {
if addr == "" {
panic("addr is empty")
}
return func(cfg *Config) {
cfg.Addr = addr
}
}
func WithTLSConfig(tlsConfig *tls.Config) Option {
return func(cfg *Config) {
cfg.TLSConfig = tlsConfig
}
}
func WithInsecure(on bool) Option {
return func(cfg *Config) {
if on {
cfg.TLSConfig = nil
} else {
cfg.TLSConfig = &tls.Config{InsecureSkipVerify: true}
}
}
}
func WithUser(user string) Option {
if user == "" {
panic("user is empty")
}
return func(cfg *Config) {
cfg.User = user
}
}
func WithPassword(password string) Option {
return func(cfg *Config) {
cfg.Password = password
}
}
func WithDatabase(database string) Option {
if database == "" {
panic("database is empty")
}
return func(cfg *Config) {
cfg.Database = database
}
}
func WithApplicationName(appName string) Option {
return func(cfg *Config) {
cfg.AppName = appName
}
}
func WithConnParams(params map[string]interface{}) Option {
return func(cfg *Config) {
cfg.ConnParams = params
}
}
func WithTimeout(timeout time.Duration) Option {
return func(cfg *Config) {
cfg.DialTimeout = timeout
cfg.ReadTimeout = timeout
cfg.WriteTimeout = timeout
}
}
func WithDialTimeout(dialTimeout time.Duration) Option {
return func(cfg *Config) {
cfg.DialTimeout = dialTimeout
}
}
func WithReadTimeout(readTimeout time.Duration) Option {
return func(cfg *Config) {
cfg.ReadTimeout = readTimeout
}
}
func WithWriteTimeout(writeTimeout time.Duration) Option {
return func(cfg *Config) {
cfg.WriteTimeout = writeTimeout
}
}
// WithResetSessionFunc configures a function that is called prior to executing
// a query on a connection that has been used before.
// If the func returns driver.ErrBadConn, the connection is discarded.
func WithResetSessionFunc(fn func(context.Context, *Conn) error) Option {
return func(cfg *Config) {
cfg.ResetSessionFunc = fn
}
}
func WithDSN(dsn string) Option {
return func(cfg *Config) {
opts, err := parseDSN(dsn)
if err != nil {
panic(err)
}
for _, opt := range opts {
opt(cfg)
}
}
}
func env(key, defValue string) string {
if s := os.Getenv(key); s != "" {
return s
}
return defValue
}
//------------------------------------------------------------------------------
func parseDSN(dsn string) ([]Option, error) {
u, err := url.Parse(dsn)
if err != nil {
return nil, err
}
q := queryOptions{q: u.Query()}
var opts []Option
switch u.Scheme {
case "postgres", "postgresql":
if u.Host != "" {
addr := u.Host
if !strings.Contains(addr, ":") {
addr += ":5432"
}
opts = append(opts, WithAddr(addr))
}
if len(u.Path) > 1 {
opts = append(opts, WithDatabase(u.Path[1:]))
}
if host := q.string("host"); host != "" {
opts = append(opts, WithAddr(host))
if host[0] == '/' {
opts = append(opts, WithNetwork("unix"))
}
}
case "unix":
if len(u.Path) == 0 {
return nil, fmt.Errorf("unix socket DSN requires a path: %s", dsn)
}
opts = append(opts, WithNetwork("unix"))
if u.Host != "" {
opts = append(opts, WithDatabase(u.Host))
}
opts = append(opts, WithAddr(u.Path))
default:
return nil, errors.New("pgdriver: invalid scheme: " + u.Scheme)
}
if u.User != nil {
opts = append(opts, WithUser(u.User.Username()))
if password, ok := u.User.Password(); ok {
opts = append(opts, WithPassword(password))
}
}
if appName := q.string("application_name"); appName != "" {
opts = append(opts, WithApplicationName(appName))
}
if sslMode, sslRootCert := q.string("sslmode"), q.string("sslrootcert"); sslMode != "" || sslRootCert != "" {
tlsConfig := &tls.Config{}
switch sslMode {
case "disable":
tlsConfig = nil
case "allow", "prefer", "":
tlsConfig.InsecureSkipVerify = true
case "require":
if sslRootCert == "" {
tlsConfig.InsecureSkipVerify = true
break
}
// For backwards compatibility reasons, in the presence of `sslrootcert`,
// `sslmode` = `require` must act as if `sslmode` = `verify-ca`. See the note at
// https://www.postgresql.org/docs/current/libpq-ssl.html#LIBQ-SSL-CERTIFICATES .
fallthrough
case "verify-ca":
// The default certificate verification will also verify the host name
// which is not the behavior of `verify-ca`. As such, we need to manually
// check the certificate chain.
// At the time of writing, tls.Config has no option for this behavior
// (verify chain, but skip server name).
// See https://github.com/golang/go/issues/21971 .
tlsConfig.InsecureSkipVerify = true
tlsConfig.VerifyPeerCertificate = func(rawCerts [][]byte, verifiedChains [][]*x509.Certificate) error {
certs := make([]*x509.Certificate, 0, len(rawCerts))
for _, rawCert := range rawCerts {
cert, err := x509.ParseCertificate(rawCert)
if err != nil {
return fmt.Errorf("pgdriver: failed to parse certificate: %w", err)
}
certs = append(certs, cert)
}
intermediates := x509.NewCertPool()
for _, cert := range certs[1:] {
intermediates.AddCert(cert)
}
_, err := certs[0].Verify(x509.VerifyOptions{
Roots: tlsConfig.RootCAs,
Intermediates: intermediates,
})
return err
}
case "verify-full":
tlsConfig.ServerName = u.Host
if host, _, err := net.SplitHostPort(u.Host); err == nil {
tlsConfig.ServerName = host
}
default:
return nil, fmt.Errorf("pgdriver: sslmode '%s' is not supported", sslMode)
}
if tlsConfig != nil && sslRootCert != "" {
rawCA, err := ioutil.ReadFile(sslRootCert)
if err != nil {
return nil, fmt.Errorf("pgdriver: failed to read root CA: %w", err)
}
certPool := x509.NewCertPool()
if !certPool.AppendCertsFromPEM(rawCA) {
return nil, fmt.Errorf("pgdriver: failed to append root CA")
}
tlsConfig.RootCAs = certPool
}
opts = append(opts, WithTLSConfig(tlsConfig))
}
if d := q.duration("timeout"); d != 0 {
opts = append(opts, WithTimeout(d))
}
if d := q.duration("dial_timeout"); d != 0 {
opts = append(opts, WithDialTimeout(d))
}
if d := q.duration("connect_timeout"); d != 0 {
opts = append(opts, WithDialTimeout(d))
}
if d := q.duration("read_timeout"); d != 0 {
opts = append(opts, WithReadTimeout(d))
}
if d := q.duration("write_timeout"); d != 0 {
opts = append(opts, WithWriteTimeout(d))
}
rem, err := q.remaining()
if err != nil {
return nil, q.err
}
if len(rem) > 0 {
params := make(map[string]interface{}, len(rem))
for k, v := range rem {
params[k] = v
}
opts = append(opts, WithConnParams(params))
}
return opts, nil
}
// verify is a method to make sure if the config is legitimate
// in the case it detects any errors, it returns with a non-nil error
// it can be extended to check other parameters
func (c *Config) verify() error {
if c.User == "" {
return errors.New("pgdriver: User option is empty (to configure, use WithUser).")
}
return nil
}
type queryOptions struct {
q url.Values
err error
}
func (o *queryOptions) string(name string) string {
vs := o.q[name]
if len(vs) == 0 {
return ""
}
delete(o.q, name) // enable detection of unknown parameters
return vs[len(vs)-1]
}
func (o *queryOptions) duration(name string) time.Duration {
s := o.string(name)
if s == "" {
return 0
}
// try plain number first
if i, err := strconv.Atoi(s); err == nil {
if i <= 0 {
// disable timeouts
return -1
}
return time.Duration(i) * time.Second
}
dur, err := time.ParseDuration(s)
if err == nil {
return dur
}
if o.err == nil {
o.err = fmt.Errorf("pgdriver: invalid %s duration: %w", name, err)
}
return 0
}
func (o *queryOptions) remaining() (map[string]string, error) {
if o.err != nil {
return nil, o.err
}
if len(o.q) == 0 {
return nil, nil
}
m := make(map[string]string, len(o.q))
for k, ss := range o.q {
m[k] = ss[len(ss)-1]
}
return m, nil
}
+249
View File
@@ -0,0 +1,249 @@
package pgdriver
import (
"bufio"
"context"
"database/sql"
"fmt"
"io"
"github.com/uptrace/bun"
)
// CopyFrom copies data from the reader to the query destination.
func CopyFrom(
ctx context.Context, conn bun.Conn, r io.Reader, query string, args ...interface{},
) (res sql.Result, err error) {
query, err = formatQueryArgs(query, args)
if err != nil {
return nil, err
}
if err := conn.Raw(func(driverConn interface{}) error {
cn := driverConn.(*Conn)
if err := writeQuery(ctx, cn, query); err != nil {
return err
}
if err := readCopyIn(ctx, cn); err != nil {
return err
}
if err := writeCopyData(ctx, cn, r); err != nil {
return err
}
if err := writeCopyDone(ctx, cn); err != nil {
return err
}
res, err = readQuery(ctx, cn)
return err
}); err != nil {
return nil, err
}
return res, nil
}
func readCopyIn(ctx context.Context, cn *Conn) error {
rd := cn.reader(ctx, -1)
var firstErr error
for {
c, msgLen, err := readMessageType(rd)
if err != nil {
return err
}
switch c {
case errorResponseMsg:
e, err := readError(rd)
if err != nil {
return err
}
if firstErr == nil {
firstErr = e
}
case readyForQueryMsg:
if err := rd.Discard(msgLen); err != nil {
return err
}
return firstErr
case copyInResponseMsg:
if err := rd.Discard(msgLen); err != nil {
return err
}
return firstErr
case noticeResponseMsg, parameterStatusMsg:
if err := rd.Discard(msgLen); err != nil {
return err
}
default:
return fmt.Errorf("pgdriver: readCopyIn: unexpected message %q", c)
}
}
}
func writeCopyData(ctx context.Context, cn *Conn, r io.Reader) error {
wb := getWriteBuffer()
defer putWriteBuffer(wb)
for {
wb.StartMessage(copyDataMsg)
if _, err := wb.ReadFrom(r); err != nil {
if err == io.EOF {
break
}
return err
}
wb.FinishMessage()
if err := cn.write(ctx, wb); err != nil {
return err
}
}
return nil
}
func writeCopyDone(ctx context.Context, cn *Conn) error {
wb := getWriteBuffer()
defer putWriteBuffer(wb)
wb.StartMessage(copyDoneMsg)
wb.FinishMessage()
return cn.write(ctx, wb)
}
//------------------------------------------------------------------------------
// CopyTo copies data from the query source to the writer.
func CopyTo(
ctx context.Context, conn bun.Conn, w io.Writer, query string, args ...interface{},
) (res sql.Result, err error) {
query, err = formatQueryArgs(query, args)
if err != nil {
return nil, err
}
if err := conn.Raw(func(driverConn interface{}) error {
cn := driverConn.(*Conn)
if err := writeQuery(ctx, cn, query); err != nil {
return err
}
if err := readCopyOut(ctx, cn); err != nil {
return err
}
res, err = readCopyData(ctx, cn, w)
return err
}); err != nil {
return nil, err
}
return res, nil
}
func readCopyOut(ctx context.Context, cn *Conn) error {
rd := cn.reader(ctx, -1)
var firstErr error
for {
c, msgLen, err := readMessageType(rd)
if err != nil {
return err
}
switch c {
case errorResponseMsg:
e, err := readError(rd)
if err != nil {
return err
}
if firstErr == nil {
firstErr = e
}
case readyForQueryMsg:
if err := rd.Discard(msgLen); err != nil {
return err
}
return firstErr
case copyOutResponseMsg:
if err := rd.Discard(msgLen); err != nil {
return err
}
return nil
case noticeResponseMsg, parameterStatusMsg:
if err := rd.Discard(msgLen); err != nil {
return err
}
default:
return fmt.Errorf("pgdriver: readCopyOut: unexpected message %q", c)
}
}
}
func readCopyData(ctx context.Context, cn *Conn, w io.Writer) (res sql.Result, err error) {
rd := cn.reader(ctx, -1)
var firstErr error
for {
c, msgLen, err := readMessageType(rd)
if err != nil {
return nil, err
}
switch c {
case errorResponseMsg:
e, err := readError(rd)
if err != nil {
return nil, err
}
if firstErr == nil {
firstErr = e
}
case copyDataMsg:
for msgLen > 0 {
b, err := rd.ReadTemp(msgLen)
if err != nil && err != bufio.ErrBufferFull {
return nil, err
}
if _, err := w.Write(b); err != nil {
if firstErr == nil {
firstErr = err
}
break
}
msgLen -= len(b)
}
case copyDoneMsg:
if err := rd.Discard(msgLen); err != nil {
return nil, err
}
case commandCompleteMsg:
tmp, err := rd.ReadTemp(msgLen)
if err != nil {
firstErr = err
break
}
r, err := parseResult(tmp)
if err != nil {
firstErr = err
} else {
res = r
}
case readyForQueryMsg:
if err := rd.Discard(msgLen); err != nil {
return nil, err
}
return res, firstErr
case noticeResponseMsg, parameterStatusMsg:
if err := rd.Discard(msgLen); err != nil {
return nil, err
}
default:
return nil, fmt.Errorf("pgdriver: readCopyData: unexpected message %q", c)
}
}
}
+600
View File
@@ -0,0 +1,600 @@
package pgdriver
import (
"bytes"
"context"
"database/sql"
"database/sql/driver"
"errors"
"fmt"
"io"
"log"
"net"
"os"
"strconv"
"sync/atomic"
"time"
)
func init() {
sql.Register("pg", NewDriver())
}
type logging interface {
Printf(ctx context.Context, format string, v ...interface{})
}
type logger struct {
log *log.Logger
}
func (l *logger) Printf(ctx context.Context, format string, v ...interface{}) {
_ = l.log.Output(2, fmt.Sprintf(format, v...))
}
var Logger logging = &logger{
log: log.New(os.Stderr, "pgdriver: ", log.LstdFlags|log.Lshortfile),
}
//------------------------------------------------------------------------------
type Driver struct {
connector *Connector
}
var _ driver.DriverContext = (*Driver)(nil)
func NewDriver() Driver {
return Driver{}
}
func (d Driver) OpenConnector(name string) (driver.Connector, error) {
opts, err := parseDSN(name)
if err != nil {
return nil, err
}
return NewConnector(opts...), nil
}
func (d Driver) Open(name string) (driver.Conn, error) {
connector, err := d.OpenConnector(name)
if err != nil {
return nil, err
}
return connector.Connect(context.TODO())
}
//------------------------------------------------------------------------------
type Connector struct {
cfg *Config
}
func NewConnector(opts ...Option) *Connector {
c := &Connector{cfg: newDefaultConfig()}
for _, opt := range opts {
opt(c.cfg)
}
return c
}
var _ driver.Connector = (*Connector)(nil)
func (c *Connector) Connect(ctx context.Context) (driver.Conn, error) {
if err := c.cfg.verify(); err != nil {
return nil, err
}
return newConn(ctx, c.cfg)
}
func (c *Connector) Driver() driver.Driver {
return Driver{connector: c}
}
func (c *Connector) Config() *Config {
return c.cfg
}
//------------------------------------------------------------------------------
type Conn struct {
cfg *Config
netConn net.Conn
rd *reader
processID int32
secretKey int32
stmtCount int
closed int32
}
func newConn(ctx context.Context, cfg *Config) (*Conn, error) {
netConn, err := cfg.Dialer(ctx, cfg.Network, cfg.Addr)
if err != nil {
return nil, err
}
cn := &Conn{
cfg: cfg,
netConn: netConn,
rd: newReader(netConn),
}
if cfg.TLSConfig != nil {
if err := enableSSL(ctx, cn, cfg.TLSConfig); err != nil {
return nil, err
}
}
if err := startup(ctx, cn); err != nil {
return nil, err
}
for k, v := range cfg.ConnParams {
if v != nil {
_, err = cn.ExecContext(ctx, fmt.Sprintf("SET %s TO $1", k), []driver.NamedValue{
{Value: v},
})
} else {
_, err = cn.ExecContext(ctx, fmt.Sprintf("SET %s TO DEFAULT", k), nil)
}
if err != nil {
return nil, err
}
}
return cn, nil
}
func (cn *Conn) reader(ctx context.Context, timeout time.Duration) *reader {
cn.setReadDeadline(ctx, timeout)
return cn.rd
}
func (cn *Conn) write(ctx context.Context, wb *writeBuffer) error {
cn.setWriteDeadline(ctx, -1)
n, err := cn.netConn.Write(wb.Bytes)
wb.Reset()
if err != nil {
if n == 0 {
Logger.Printf(ctx, "pgdriver: Conn.Write failed (zero-length): %s", err)
return driver.ErrBadConn
}
return err
}
return nil
}
var _ driver.Conn = (*Conn)(nil)
func (cn *Conn) Prepare(query string) (driver.Stmt, error) {
if cn.isClosed() {
return nil, driver.ErrBadConn
}
ctx := context.TODO()
name := fmt.Sprintf("pgdriver-%d", cn.stmtCount)
cn.stmtCount++
if err := writeParseDescribeSync(ctx, cn, name, query); err != nil {
return nil, err
}
rowDesc, err := readParseDescribeSync(ctx, cn)
if err != nil {
return nil, err
}
return newStmt(cn, name, rowDesc), nil
}
func (cn *Conn) Close() error {
if !atomic.CompareAndSwapInt32(&cn.closed, 0, 1) {
return nil
}
return cn.netConn.Close()
}
func (cn *Conn) isClosed() bool {
return atomic.LoadInt32(&cn.closed) == 1
}
func (cn *Conn) Begin() (driver.Tx, error) {
return cn.BeginTx(context.Background(), driver.TxOptions{})
}
var _ driver.ConnBeginTx = (*Conn)(nil)
func (cn *Conn) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) {
// No need to check if the conn is closed. ExecContext below handles that.
if sql.IsolationLevel(opts.Isolation) != sql.LevelDefault {
return nil, errors.New("pgdriver: custom IsolationLevel is not supported")
}
if opts.ReadOnly {
return nil, errors.New("pgdriver: ReadOnly transactions are not supported")
}
if _, err := cn.ExecContext(ctx, "BEGIN", nil); err != nil {
return nil, err
}
return tx{cn: cn}, nil
}
var _ driver.ExecerContext = (*Conn)(nil)
func (cn *Conn) ExecContext(
ctx context.Context, query string, args []driver.NamedValue,
) (driver.Result, error) {
if cn.isClosed() {
return nil, driver.ErrBadConn
}
res, err := cn.exec(ctx, query, args)
if err != nil {
return nil, cn.checkBadConn(err)
}
return res, nil
}
func (cn *Conn) exec(
ctx context.Context, query string, args []driver.NamedValue,
) (driver.Result, error) {
query, err := formatQuery(query, args)
if err != nil {
return nil, err
}
if err := writeQuery(ctx, cn, query); err != nil {
return nil, err
}
return readQuery(ctx, cn)
}
var _ driver.QueryerContext = (*Conn)(nil)
func (cn *Conn) QueryContext(
ctx context.Context, query string, args []driver.NamedValue,
) (driver.Rows, error) {
if cn.isClosed() {
return nil, driver.ErrBadConn
}
rows, err := cn.query(ctx, query, args)
if err != nil {
return nil, cn.checkBadConn(err)
}
return rows, nil
}
func (cn *Conn) query(
ctx context.Context, query string, args []driver.NamedValue,
) (driver.Rows, error) {
query, err := formatQuery(query, args)
if err != nil {
return nil, err
}
if err := writeQuery(ctx, cn, query); err != nil {
return nil, err
}
return readQueryData(ctx, cn)
}
var _ driver.Pinger = (*Conn)(nil)
func (cn *Conn) Ping(ctx context.Context) error {
_, err := cn.ExecContext(ctx, "SELECT 1", nil)
return err
}
func (cn *Conn) setReadDeadline(ctx context.Context, timeout time.Duration) {
if timeout == -1 {
timeout = cn.cfg.ReadTimeout
}
_ = cn.netConn.SetReadDeadline(cn.deadline(ctx, timeout))
}
func (cn *Conn) setWriteDeadline(ctx context.Context, timeout time.Duration) {
if timeout == -1 {
timeout = cn.cfg.WriteTimeout
}
_ = cn.netConn.SetWriteDeadline(cn.deadline(ctx, timeout))
}
func (cn *Conn) deadline(ctx context.Context, timeout time.Duration) time.Time {
deadline, ok := ctx.Deadline()
if !ok {
if timeout == 0 {
return time.Time{}
}
return time.Now().Add(timeout)
}
if timeout == 0 {
return deadline
}
if tm := time.Now().Add(timeout); tm.Before(deadline) {
return tm
}
return deadline
}
var _ driver.Validator = (*Conn)(nil)
func (cn *Conn) IsValid() bool {
return !cn.isClosed()
}
var _ driver.SessionResetter = (*Conn)(nil)
func (cn *Conn) ResetSession(ctx context.Context) error {
if cn.isClosed() {
return driver.ErrBadConn
}
if cn.cfg.ResetSessionFunc != nil {
return cn.cfg.ResetSessionFunc(ctx, cn)
}
return nil
}
func (cn *Conn) checkBadConn(err error) error {
if isBadConn(err, false) {
// Close and return driver.ErrBadConn next time the conn is used.
_ = cn.Close()
}
// Always return the original error.
return err
}
func (cn *Conn) Conn() net.Conn { return cn.netConn }
//------------------------------------------------------------------------------
type rows struct {
cn *Conn
rowDesc *rowDescription
reusable bool
closed bool
}
var _ driver.Rows = (*rows)(nil)
func newRows(cn *Conn, rowDesc *rowDescription, reusable bool) *rows {
return &rows{
cn: cn,
rowDesc: rowDesc,
reusable: reusable,
}
}
func (r *rows) Columns() []string {
if r.closed || r.rowDesc == nil {
return nil
}
return r.rowDesc.names
}
func (r *rows) Close() error {
if r.closed {
return nil
}
defer r.close()
for {
switch err := r.Next(nil); err {
case nil, io.EOF:
return nil
default: // unexpected error
_ = r.cn.Close()
return err
}
}
}
func (r *rows) close() {
r.closed = true
if r.rowDesc != nil {
if r.reusable {
rowDescPool.Put(r.rowDesc)
}
r.rowDesc = nil
}
}
func (r *rows) Next(dest []driver.Value) error {
if r.closed {
return io.EOF
}
eof, err := r.next(dest)
if err == io.EOF {
return io.ErrUnexpectedEOF
} else if err != nil {
return err
}
if eof {
return io.EOF
}
return nil
}
func (r *rows) next(dest []driver.Value) (eof bool, _ error) {
rd := r.cn.reader(context.TODO(), -1)
var firstErr error
for {
c, msgLen, err := readMessageType(rd)
if err != nil {
return false, err
}
switch c {
case dataRowMsg:
return false, r.readDataRow(rd, dest)
case commandCompleteMsg:
if err := rd.Discard(msgLen); err != nil {
return false, err
}
case readyForQueryMsg:
r.close()
if err := rd.Discard(msgLen); err != nil {
return false, err
}
if firstErr != nil {
return false, firstErr
}
return true, nil
case parameterStatusMsg, noticeResponseMsg:
if err := rd.Discard(msgLen); err != nil {
return false, err
}
case errorResponseMsg:
e, err := readError(rd)
if err != nil {
return false, err
}
if firstErr == nil {
firstErr = e
}
default:
return false, fmt.Errorf("pgdriver: Next: unexpected message %q", c)
}
}
}
func (r *rows) readDataRow(rd *reader, dest []driver.Value) error {
numCol, err := readInt16(rd)
if err != nil {
return err
}
if len(dest) != int(numCol) {
return fmt.Errorf("pgdriver: query returned %d columns, but Scan dest has %d items",
numCol, len(dest))
}
for colIdx := int16(0); colIdx < numCol; colIdx++ {
dataLen, err := readInt32(rd)
if err != nil {
return err
}
value, err := readColumnValue(rd, r.rowDesc.types[colIdx], int(dataLen))
if err != nil {
return err
}
if dest != nil {
dest[colIdx] = value
}
}
return nil
}
//------------------------------------------------------------------------------
func parseResult(b []byte) (driver.RowsAffected, error) {
i := bytes.LastIndexByte(b, ' ')
if i == -1 {
return 0, nil
}
b = b[i+1 : len(b)-1]
affected, err := strconv.ParseUint(bytesToString(b), 10, 64)
if err != nil {
return 0, nil
}
return driver.RowsAffected(affected), nil
}
//------------------------------------------------------------------------------
type tx struct {
cn *Conn
}
var _ driver.Tx = (*tx)(nil)
func (tx tx) Commit() error {
_, err := tx.cn.ExecContext(context.Background(), "COMMIT", nil)
return err
}
func (tx tx) Rollback() error {
_, err := tx.cn.ExecContext(context.Background(), "ROLLBACK", nil)
return err
}
//------------------------------------------------------------------------------
type stmt struct {
cn *Conn
name string
rowDesc *rowDescription
}
var (
_ driver.Stmt = (*stmt)(nil)
_ driver.StmtExecContext = (*stmt)(nil)
_ driver.StmtQueryContext = (*stmt)(nil)
)
func newStmt(cn *Conn, name string, rowDesc *rowDescription) *stmt {
return &stmt{
cn: cn,
name: name,
rowDesc: rowDesc,
}
}
func (stmt *stmt) Close() error {
if stmt.rowDesc != nil {
rowDescPool.Put(stmt.rowDesc)
stmt.rowDesc = nil
}
ctx := context.TODO()
if err := writeCloseStmt(ctx, stmt.cn, stmt.name); err != nil {
return err
}
if err := readCloseStmtComplete(ctx, stmt.cn); err != nil {
return err
}
return nil
}
func (stmt *stmt) NumInput() int {
if stmt.rowDesc == nil {
return -1
}
return int(stmt.rowDesc.numInput)
}
func (stmt *stmt) Exec(args []driver.Value) (driver.Result, error) {
panic("not implemented")
}
func (stmt *stmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) {
if err := writeBindExecute(ctx, stmt.cn, stmt.name, args); err != nil {
return nil, err
}
return readExtQuery(ctx, stmt.cn)
}
func (stmt *stmt) Query(args []driver.Value) (driver.Rows, error) {
panic("not implemented")
}
func (stmt *stmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) {
if err := writeBindExecute(ctx, stmt.cn, stmt.name, args); err != nil {
return nil, err
}
return readExtQueryData(ctx, stmt.cn, stmt.rowDesc)
}
+75
View File
@@ -0,0 +1,75 @@
package pgdriver
import (
"database/sql/driver"
"fmt"
"net"
)
// Error represents an error returned by PostgreSQL server
// using PostgreSQL ErrorResponse protocol.
//
// https://www.postgresql.org/docs/current/static/protocol-message-formats.html
type Error struct {
m map[byte]string
}
// Field returns a string value associated with an error field.
//
// https://www.postgresql.org/docs/current/static/protocol-error-fields.html
func (err Error) Field(k byte) string {
return err.m[k]
}
// IntegrityViolation reports whether the error is a part of
// Integrity Constraint Violation class of errors.
//
// https://www.postgresql.org/docs/current/static/errcodes-appendix.html
func (err Error) IntegrityViolation() bool {
switch err.Field('C') {
case "23000", "23001", "23502", "23503", "23505", "23514", "23P01":
return true
default:
return false
}
}
// StatementTimeout reports whether the error is a statement timeout error.
func (err Error) StatementTimeout() bool {
return err.Field('C') == "57014"
}
func (err Error) Error() string {
return fmt.Sprintf("%s: %s (SQLSTATE=%s)",
err.Field('S'), err.Field('M'), err.Field('C'))
}
func isBadConn(err error, allowTimeout bool) bool {
switch err {
case nil:
return false
case driver.ErrBadConn:
return true
}
if err, ok := err.(Error); ok {
switch err.Field('V') {
case "FATAL", "PANIC":
return true
}
switch err.Field('C') {
case "25P02", // current transaction is aborted
"57014": // canceling statement due to user request
return true
}
return false
}
if allowTimeout {
if err, ok := err.(net.Error); ok && err.Timeout() {
return !err.Temporary()
}
}
return true
}
+199
View File
@@ -0,0 +1,199 @@
package pgdriver
import (
"database/sql/driver"
"encoding/hex"
"fmt"
"math"
"strconv"
"time"
"unicode/utf8"
)
func formatQueryArgs(query string, args []interface{}) (string, error) {
namedArgs := make([]driver.NamedValue, len(args))
for i, arg := range args {
namedArgs[i] = driver.NamedValue{Value: arg}
}
return formatQuery(query, namedArgs)
}
func formatQuery(query string, args []driver.NamedValue) (string, error) {
if len(args) == 0 {
return query, nil
}
dst := make([]byte, 0, 2*len(query))
p := newParser(query)
for p.Valid() {
switch c := p.Next(); c {
case '$':
if i, ok := p.Number(); ok {
if i < 1 {
return "", fmt.Errorf("pgdriver: got $%d, but the minimal arg index is 1", i)
}
if i > len(args) {
return "", fmt.Errorf("pgdriver: got %d args, wanted %d", len(args), i)
}
var err error
dst, err = appendArg(dst, args[i-1].Value)
if err != nil {
return "", err
}
} else {
dst = append(dst, '$')
}
case '\'':
if b, ok := p.QuotedString(); ok {
dst = append(dst, b...)
} else {
dst = append(dst, '\'')
}
default:
dst = append(dst, c)
}
}
return bytesToString(dst), nil
}
func appendArg(b []byte, v interface{}) ([]byte, error) {
switch v := v.(type) {
case nil:
return append(b, "NULL"...), nil
case int64:
return strconv.AppendInt(b, v, 10), nil
case float64:
switch {
case math.IsNaN(v):
return append(b, "'NaN'"...), nil
case math.IsInf(v, 1):
return append(b, "'Infinity'"...), nil
case math.IsInf(v, -1):
return append(b, "'-Infinity'"...), nil
default:
return strconv.AppendFloat(b, v, 'f', -1, 64), nil
}
case bool:
if v {
return append(b, "TRUE"...), nil
}
return append(b, "FALSE"...), nil
case []byte:
if v == nil {
return append(b, "NULL"...), nil
}
b = append(b, `'\x`...)
s := len(b)
b = append(b, make([]byte, hex.EncodedLen(len(v)))...)
hex.Encode(b[s:], v)
b = append(b, "'"...)
return b, nil
case string:
b = append(b, '\'')
for _, r := range v {
if r == '\000' {
continue
}
if r == '\'' {
b = append(b, '\'', '\'')
continue
}
if r < utf8.RuneSelf {
b = append(b, byte(r))
continue
}
l := len(b)
if cap(b)-l < utf8.UTFMax {
b = append(b, make([]byte, utf8.UTFMax)...)
}
n := utf8.EncodeRune(b[l:l+utf8.UTFMax], r)
b = b[:l+n]
}
b = append(b, '\'')
return b, nil
case time.Time:
if v.IsZero() {
return append(b, "NULL"...), nil
}
return v.UTC().AppendFormat(b, "'2006-01-02 15:04:05.999999-07:00'"), nil
default:
return nil, fmt.Errorf("pgdriver: unexpected arg: %T", v)
}
}
type parser struct {
b []byte
i int
}
func newParser(s string) *parser {
return &parser{
b: stringToBytes(s),
}
}
func (p *parser) Valid() bool {
return p.i < len(p.b)
}
func (p *parser) Next() byte {
c := p.b[p.i]
p.i++
return c
}
func (p *parser) Number() (int, bool) {
start := p.i
end := len(p.b)
for i := p.i; i < len(p.b); i++ {
c := p.b[i]
if !isNum(c) {
end = i
break
}
}
p.i = end
b := p.b[start:end]
n, err := strconv.Atoi(bytesToString(b))
if err != nil {
return 0, false
}
return n, true
}
func (p *parser) QuotedString() ([]byte, bool) {
start := p.i - 1
end := len(p.b)
var c byte
for i := p.i; i < len(p.b); i++ {
next := p.b[i]
if c == '\'' && next != '\'' {
end = i
break
}
c = next
}
p.i = end
b := p.b[start:end]
return b, true
}
func isNum(c byte) bool {
return c >= '0' && c <= '9'
}
+380
View File
@@ -0,0 +1,380 @@
package pgdriver
import (
"context"
"errors"
"strconv"
"sync"
"time"
"github.com/uptrace/bun"
)
const pingChannel = "bun:ping"
var (
errListenerClosed = errors.New("bun: listener is closed")
errPingTimeout = errors.New("bun: ping timeout")
)
// Notify sends a notification on the channel using `NOTIFY` command.
func Notify(ctx context.Context, db *bun.DB, channel, payload string) error {
_, err := db.ExecContext(ctx, "NOTIFY ?, ?", bun.Ident(channel), payload)
return err
}
type Listener struct {
db *bun.DB
driver *Connector
channels []string
mu sync.Mutex
cn *Conn
closed bool
exit chan struct{}
}
func NewListener(db *bun.DB) *Listener {
return &Listener{
db: db,
driver: db.Driver().(Driver).connector,
exit: make(chan struct{}),
}
}
// Close closes the listener, releasing any open resources.
func (ln *Listener) Close() error {
return ln.withLock(func() error {
if ln.closed {
return errListenerClosed
}
ln.closed = true
close(ln.exit)
return ln.closeConn(errListenerClosed)
})
}
func (ln *Listener) withLock(fn func() error) error {
ln.mu.Lock()
defer ln.mu.Unlock()
return fn()
}
func (ln *Listener) conn(ctx context.Context) (*Conn, error) {
if ln.closed {
return nil, errListenerClosed
}
if ln.cn != nil {
return ln.cn, nil
}
cn, err := ln._conn(ctx)
if err != nil {
return nil, err
}
ln.cn = cn
return cn, nil
}
func (ln *Listener) _conn(ctx context.Context) (*Conn, error) {
driverConn, err := ln.driver.Connect(ctx)
if err != nil {
return nil, err
}
cn := driverConn.(*Conn)
if len(ln.channels) > 0 {
err := ln.listen(ctx, cn, ln.channels...)
if err != nil {
_ = cn.Close()
return nil, err
}
}
return cn, nil
}
func (ln *Listener) checkConn(ctx context.Context, cn *Conn, err error, allowTimeout bool) {
_ = ln.withLock(func() error {
if ln.closed || ln.cn != cn {
return nil
}
if isBadConn(err, allowTimeout) {
ln.reconnect(ctx, err)
}
return nil
})
}
func (ln *Listener) reconnect(ctx context.Context, reason error) {
if ln.cn != nil {
Logger.Printf(ctx, "bun: discarding bad listener connection: %s", reason)
_ = ln.closeConn(reason)
}
_, _ = ln.conn(ctx)
}
func (ln *Listener) closeConn(reason error) error {
if ln.cn == nil {
return nil
}
err := ln.cn.Close()
ln.cn = nil
return err
}
// Listen starts listening for notifications on channels.
func (ln *Listener) Listen(ctx context.Context, channels ...string) error {
var cn *Conn
if err := ln.withLock(func() error {
ln.channels = appendIfNotExists(ln.channels, channels...)
var err error
cn, err = ln.conn(ctx)
return err
}); err != nil {
return err
}
if err := ln.listen(ctx, cn, channels...); err != nil {
ln.checkConn(ctx, cn, err, false)
return err
}
return nil
}
func (ln *Listener) listen(ctx context.Context, cn *Conn, channels ...string) error {
for _, channel := range channels {
if err := writeQuery(ctx, cn, "LISTEN "+strconv.Quote(channel)); err != nil {
return err
}
}
return nil
}
// Unlisten stops listening for notifications on channels.
func (ln *Listener) Unlisten(ctx context.Context, channels ...string) error {
var cn *Conn
if err := ln.withLock(func() error {
ln.channels = removeIfExists(ln.channels, channels...)
var err error
cn, err = ln.conn(ctx)
return err
}); err != nil {
return err
}
if err := ln.unlisten(ctx, cn, channels...); err != nil {
ln.checkConn(ctx, cn, err, false)
return err
}
return nil
}
func (ln *Listener) unlisten(ctx context.Context, cn *Conn, channels ...string) error {
for _, channel := range channels {
if err := writeQuery(ctx, cn, "UNLISTEN "+strconv.Quote(channel)); err != nil {
return err
}
}
return nil
}
// Receive indefinitely waits for a notification. This is low-level API
// and in most cases Channel should be used instead.
func (ln *Listener) Receive(ctx context.Context) (channel string, payload string, err error) {
return ln.ReceiveTimeout(ctx, 0)
}
// ReceiveTimeout waits for a notification until timeout is reached.
// This is low-level API and in most cases Channel should be used instead.
func (ln *Listener) ReceiveTimeout(
ctx context.Context, timeout time.Duration,
) (channel, payload string, err error) {
var cn *Conn
if err := ln.withLock(func() error {
var err error
cn, err = ln.conn(ctx)
return err
}); err != nil {
return "", "", err
}
rd := cn.reader(ctx, timeout)
channel, payload, err = readNotification(ctx, rd)
if err != nil {
ln.checkConn(ctx, cn, err, timeout > 0)
return "", "", err
}
return channel, payload, nil
}
// Channel returns a channel for concurrently receiving notifications.
// It periodically sends Ping notification to test connection health.
//
// The channel is closed with Listener. Receive* APIs can not be used
// after channel is created.
func (ln *Listener) Channel(opts ...ChannelOption) <-chan Notification {
return newChannel(ln, opts).ch
}
//------------------------------------------------------------------------------
// Notification received with LISTEN command.
type Notification struct {
Channel string
Payload string
}
type ChannelOption func(c *channel)
func WithChannelSize(size int) ChannelOption {
return func(c *channel) {
c.size = size
}
}
type channel struct {
ctx context.Context
ln *Listener
size int
pingTimeout time.Duration
ch chan Notification
pingCh chan struct{}
}
func newChannel(ln *Listener, opts []ChannelOption) *channel {
c := &channel{
ctx: context.TODO(),
ln: ln,
size: 1000,
pingTimeout: 5 * time.Second,
}
for _, opt := range opts {
opt(c)
}
c.ch = make(chan Notification, c.size)
c.pingCh = make(chan struct{}, 1)
_ = c.ln.Listen(c.ctx, pingChannel)
go c.startReceive()
go c.startPing()
return c
}
func (c *channel) startReceive() {
var errCount int
for {
channel, payload, err := c.ln.Receive(c.ctx)
if err != nil {
if err == errListenerClosed {
close(c.ch)
return
}
if errCount > 0 {
time.Sleep(500 * time.Millisecond)
}
errCount++
continue
}
errCount = 0
// Any notification is as good as a ping.
select {
case c.pingCh <- struct{}{}:
default:
}
switch channel {
case pingChannel:
// ignore
default:
select {
case c.ch <- Notification{channel, payload}:
default:
Logger.Printf(c.ctx, "pgdriver: Listener buffer is full (message is dropped)")
}
}
}
}
func (c *channel) startPing() {
timer := time.NewTimer(time.Minute)
timer.Stop()
healthy := true
for {
timer.Reset(c.pingTimeout)
select {
case <-c.pingCh:
healthy = true
if !timer.Stop() {
<-timer.C
}
case <-timer.C:
pingErr := c.ping(c.ctx)
if healthy {
healthy = false
} else {
if pingErr == nil {
pingErr = errPingTimeout
}
_ = c.ln.withLock(func() error {
c.ln.reconnect(c.ctx, pingErr)
return nil
})
}
case <-c.ln.exit:
return
}
}
}
func (c *channel) ping(ctx context.Context) error {
_, err := c.ln.db.ExecContext(ctx, "NOTIFY "+strconv.Quote(pingChannel))
return err
}
func appendIfNotExists(ss []string, es ...string) []string {
loop:
for _, e := range es {
for _, s := range ss {
if s == e {
continue loop
}
}
ss = append(ss, e)
}
return ss
}
func removeIfExists(ss []string, es ...string) []string {
for _, e := range es {
for i, s := range ss {
if s == e {
last := len(ss) - 1
ss[i] = ss[last]
ss = ss[:last]
break
}
}
}
return ss
}
File diff suppressed because it is too large Load Diff
+11
View File
@@ -0,0 +1,11 @@
// +build appengine
package internal
func bytesToString(b []byte) string {
return string(b)
}
func stringToBytes(s string) []byte {
return []byte(s)
}
+19
View File
@@ -0,0 +1,19 @@
// +build !appengine
package pgdriver
import "unsafe"
func bytesToString(b []byte) string {
return *(*string)(unsafe.Pointer(&b))
}
//nolint:deadcode,unused
func stringToBytes(s string) []byte {
return *(*[]byte)(unsafe.Pointer(
&struct {
string
Cap int
}{s, len(s)},
))
}
+112
View File
@@ -0,0 +1,112 @@
package pgdriver
import (
"encoding/binary"
"io"
"sync"
)
var wbPool = sync.Pool{
New: func() interface{} {
return newWriteBuffer()
},
}
func getWriteBuffer() *writeBuffer {
wb := wbPool.Get().(*writeBuffer)
return wb
}
func putWriteBuffer(wb *writeBuffer) {
wb.Reset()
wbPool.Put(wb)
}
type writeBuffer struct {
Bytes []byte
msgStart int
paramStart int
}
func newWriteBuffer() *writeBuffer {
return &writeBuffer{
Bytes: make([]byte, 0, 1024),
}
}
func (b *writeBuffer) Reset() {
b.Bytes = b.Bytes[:0]
}
func (b *writeBuffer) StartMessage(c byte) {
if c == 0 {
b.msgStart = len(b.Bytes)
b.Bytes = append(b.Bytes, 0, 0, 0, 0)
} else {
b.msgStart = len(b.Bytes) + 1
b.Bytes = append(b.Bytes, c, 0, 0, 0, 0)
}
}
func (b *writeBuffer) FinishMessage() {
binary.BigEndian.PutUint32(
b.Bytes[b.msgStart:], uint32(len(b.Bytes)-b.msgStart))
}
func (b *writeBuffer) Query() []byte {
return b.Bytes[b.msgStart+4 : len(b.Bytes)-1]
}
func (b *writeBuffer) StartParam() {
b.paramStart = len(b.Bytes)
b.Bytes = append(b.Bytes, 0, 0, 0, 0)
}
func (b *writeBuffer) FinishParam() {
binary.BigEndian.PutUint32(
b.Bytes[b.paramStart:], uint32(len(b.Bytes)-b.paramStart-4))
}
var nullParamLength = int32(-1)
func (b *writeBuffer) FinishNullParam() {
binary.BigEndian.PutUint32(
b.Bytes[b.paramStart:], uint32(nullParamLength))
}
func (b *writeBuffer) Write(data []byte) (int, error) {
b.Bytes = append(b.Bytes, data...)
return len(data), nil
}
func (b *writeBuffer) WriteInt16(num int16) {
b.Bytes = append(b.Bytes, 0, 0)
binary.BigEndian.PutUint16(b.Bytes[len(b.Bytes)-2:], uint16(num))
}
func (b *writeBuffer) WriteInt32(num int32) {
b.Bytes = append(b.Bytes, 0, 0, 0, 0)
binary.BigEndian.PutUint32(b.Bytes[len(b.Bytes)-4:], uint32(num))
}
func (b *writeBuffer) WriteString(s string) {
b.Bytes = append(b.Bytes, s...)
b.Bytes = append(b.Bytes, 0)
}
func (b *writeBuffer) WriteBytes(data []byte) {
b.Bytes = append(b.Bytes, data...)
b.Bytes = append(b.Bytes, 0)
}
func (b *writeBuffer) WriteByte(c byte) error {
b.Bytes = append(b.Bytes, c)
return nil
}
func (b *writeBuffer) ReadFrom(r io.Reader) (int64, error) {
n, err := r.Read(b.Bytes[len(b.Bytes):cap(b.Bytes)])
b.Bytes = b.Bytes[:len(b.Bytes)+n]
return int64(n), err
}