Files
token_thief/db/clickhouse.go
T

393 lines
9.1 KiB
Go

package db
import (
"context"
"errors"
"fmt"
"log"
"net"
"net/url"
"strconv"
"strings"
"sync"
"time"
"github.com/ClickHouse/clickhouse-go/v2"
)
// Pool 封装 ClickHouse 连接并维护健康状态,DB 故障时不阻塞调用方。
type Pool struct {
dsn string
reconnectInterval time.Duration
open func(*clickhouse.Options) (clickhouse.Conn, error)
mu sync.RWMutex
conn clickhouse.Conn
generation uint64
healthy bool
closed bool
leases map[uint64]*connectionLease
leaseChanged chan struct{}
closing int
}
type connectionLease struct {
conn clickhouse.Conn
active int
retired bool
closed bool
}
var errPoolClosed = errors.New("db pool is closed")
// NewPool 创建 Pool。即使首次连接失败也返回非 nil 实例,后台会持续重试。
func NewPool(ctx context.Context, dsn string, reconnectInterval time.Duration) *Pool {
p := &Pool{dsn: dsn, reconnectInterval: reconnectInterval}
if err := p.connect(ctx); err != nil {
log.Printf("[db] initial connect failed: %v (service continues without DB)", err)
}
go p.watch(ctx)
return p
}
func (p *Pool) connect(ctx context.Context) error {
p.mu.RLock()
closed := p.closed
p.mu.RUnlock()
if closed {
return errPoolClosed
}
cctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
opts, err := ClickHouseOptions(p.dsn)
if err != nil {
return err
}
open := p.open
if open == nil {
open = clickhouse.Open
}
conn, err := open(opts)
if err != nil {
return err
}
if err := conn.Ping(cctx); err != nil {
_ = conn.Close()
return err
}
// 先 migrate,再 swap:避免新连接 migrate 失败时取代掉旧的可用连接。
mctx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
if err := migrate(mctx, conn); err != nil {
_ = conn.Close()
log.Printf("[db] migrate failed: %v", err)
return err
}
p.mu.Lock()
if p.closed {
p.mu.Unlock()
_ = conn.Close()
return errPoolClosed
}
old := p.retireCurrentLocked()
p.conn = conn
p.generation++
p.ensureLeaseLocked(p.generation, conn)
p.healthy = true
p.mu.Unlock()
if old != nil {
_ = old.Close()
p.finishClosing()
}
log.Printf("[db] connected and migrated")
return nil
}
// ClickHouseOptions converts the supported CLICKHOUSE_URL subset into driver options.
func ClickHouseOptions(raw string) (*clickhouse.Options, error) {
if !strings.Contains(raw, "://") {
if err := validateHostPort(raw); err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
return &clickhouse.Options{Addr: []string{raw}}, nil
}
u, err := url.Parse(raw)
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
if u.Scheme != "clickhouse" && u.Scheme != "clickhouses" {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: unsupported scheme %q", u.Scheme)
}
if u.Fragment != "" {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: fragment is not allowed")
}
if err := validateHostPort(u.Host); err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
query, err := url.ParseQuery(u.RawQuery)
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: query: %w", err)
}
for key, values := range query {
switch key {
case "secure", "skip_verify", "compress":
default:
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: unsupported parameter %q", key)
}
if len(values) != 1 {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: parameter %q must occur once", key)
}
}
wantSecure := u.Scheme == "clickhouses"
if value, ok := query["secure"]; ok {
secure, err := strconv.ParseBool(value[0])
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: secure: %w", err)
}
if secure != wantSecure {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: secure conflicts with %s", u.Scheme)
}
}
if _, ok := query["skip_verify"]; ok && !wantSecure {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: skip_verify requires clickhouses")
}
if value, ok := query["skip_verify"]; ok {
if _, err := strconv.ParseBool(value[0]); err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: skip_verify: %w", err)
}
}
if value, ok := query["compress"]; ok {
switch value[0] {
case "true", "false", "none", "zstd", "lz4", "lz4hc", "gzip", "deflate", "br":
default:
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: unsupported compression %q", value[0])
}
}
// Let the driver decode userinfo/path and interpret compression and TLS values.
if wantSecure {
query.Set("secure", "true")
u.RawQuery = query.Encode()
}
opts, err := clickhouse.ParseDSN(u.String())
if err != nil {
return nil, fmt.Errorf("invalid CLICKHOUSE_URL: %w", err)
}
return opts, nil
}
func validateHostPort(address string) error {
if address == "" {
return fmt.Errorf("missing host and port")
}
host, port, err := net.SplitHostPort(address)
if err != nil {
return fmt.Errorf("expected host:port: %w", err)
}
if host == "" {
return fmt.Errorf("missing host")
}
n, err := strconv.ParseUint(port, 10, 16)
if err != nil || n == 0 {
return fmt.Errorf("invalid port %q", port)
}
return nil
}
// watch 定期探活;不健康时尝试重连。
func (p *Pool) watch(ctx context.Context) {
t := time.NewTicker(p.reconnectInterval)
defer t.Stop()
for {
select {
case <-ctx.Done():
return
case <-t.C:
if p.Healthy() {
if conn, generation, release := p.Acquire(); conn != nil {
pctx, cancel := context.WithTimeout(ctx, 3*time.Second)
if err := conn.Ping(pctx); err != nil {
log.Printf("[db] ping failed, marking unhealthy: %v", err)
p.MarkUnhealthyGeneration(generation)
}
cancel()
release()
}
continue
}
if err := p.connect(ctx); err != nil {
log.Printf("[db] reconnect failed: %v", err)
}
}
}
}
// Healthy 报告连接是否可用。
func (p *Pool) Healthy() bool {
p.mu.RLock()
defer p.mu.RUnlock()
return p.healthy
}
// MarkUnhealthy 由调用方在写入失败后调用。
func (p *Pool) MarkUnhealthy() {
p.mu.Lock()
p.healthy = false
p.mu.Unlock()
}
// MarkUnhealthyGeneration only changes health when the failed connection is still current.
func (p *Pool) MarkUnhealthyGeneration(generation uint64) {
p.mu.Lock()
if p.generation == generation {
p.healthy = false
}
p.mu.Unlock()
}
func (p *Pool) Generation() uint64 {
p.mu.RLock()
defer p.mu.RUnlock()
return p.generation
}
// Get 返回当前连接,可能为 nil。
func (p *Pool) Get() clickhouse.Conn {
p.mu.RLock()
defer p.mu.RUnlock()
return p.conn
}
func (p *Pool) GetWithGeneration() (clickhouse.Conn, uint64) {
p.mu.RLock()
defer p.mu.RUnlock()
return p.conn, p.generation
}
// Acquire pins the current connection until release is called.
func (p *Pool) Acquire() (clickhouse.Conn, uint64, func()) {
p.mu.Lock()
if p.closed || p.conn == nil {
p.mu.Unlock()
return nil, p.generation, nil
}
generation := p.generation
lease := p.ensureLeaseLocked(generation, p.conn)
lease.active++
p.mu.Unlock()
var once sync.Once
return lease.conn, generation, func() {
once.Do(func() { p.release(generation) })
}
}
func (p *Pool) ensureLeaseLocked(generation uint64, conn clickhouse.Conn) *connectionLease {
if p.leases == nil {
p.leases = make(map[uint64]*connectionLease)
}
lease := p.leases[generation]
if lease == nil {
lease = &connectionLease{conn: conn}
p.leases[generation] = lease
}
return lease
}
func (p *Pool) notifyLeaseChangedLocked() {
if p.leaseChanged != nil {
close(p.leaseChanged)
}
p.leaseChanged = make(chan struct{})
}
func (p *Pool) retireCurrentLocked() clickhouse.Conn {
if p.conn == nil {
return nil
}
lease := p.ensureLeaseLocked(p.generation, p.conn)
lease.retired = true
if lease.active != 0 || lease.closed {
return nil
}
lease.closed = true
delete(p.leases, p.generation)
p.closing++
return lease.conn
}
func (p *Pool) finishClosing() {
p.mu.Lock()
p.closing--
p.notifyLeaseChangedLocked()
p.mu.Unlock()
}
func (p *Pool) release(generation uint64) {
p.mu.Lock()
lease := p.leases[generation]
if lease == nil {
p.mu.Unlock()
return
}
lease.active--
var conn clickhouse.Conn
if lease.active == 0 && lease.retired && !lease.closed {
lease.closed = true
conn = lease.conn
p.closing++
}
p.mu.Unlock()
if conn != nil {
_ = conn.Close()
p.mu.Lock()
delete(p.leases, generation)
p.closing--
p.notifyLeaseChangedLocked()
p.mu.Unlock()
}
}
// Close starts closing the underlying connection and waits within ctx.
func (p *Pool) Close(ctx context.Context) error {
p.mu.Lock()
p.closed = true
conn := p.retireCurrentLocked()
p.conn = nil
p.healthy = false
p.mu.Unlock()
var closeErr error
if conn != nil {
done := make(chan error, 1)
go func() {
done <- conn.Close()
p.finishClosing()
}()
select {
case closeErr = <-done:
case <-ctx.Done():
return ctx.Err()
}
}
for {
p.mu.Lock()
if len(p.leases) == 0 && p.closing == 0 {
p.mu.Unlock()
return closeErr
}
if p.leaseChanged == nil {
p.leaseChanged = make(chan struct{})
}
changed := p.leaseChanged
p.mu.Unlock()
select {
case <-changed:
case <-ctx.Done():
return ctx.Err()
}
}
}