Use atomic pointer in go 1.19 (#446)

* use atomic pointer

* golang version

* golang version

* go1.19

* linters

* Bump golangci-lint

* linters

* linters

* linters after merge

* generic logger

* generic logger

* logger

* logger

* linters

* bump toml

* linters1

* linters

* linters

* linter

* linter

* linters

* linters

* linters
This commit is contained in:
Evgeny Danilenko 2022-08-09 22:11:09 +03:00 committed by GitHub
parent ac559bcd16
commit e699254142
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
21 changed files with 175 additions and 100 deletions

View file

@ -29,7 +29,7 @@ jobs:
- uses: actions/setup-go@v3 - uses: actions/setup-go@v3
with: with:
go-version: 1.18.x go-version: 1.19.x
- name: Install dependencies on Linux - name: Install dependencies on Linux
if: runner.os == 'Linux' if: runner.os == 'Linux'

View file

@ -50,6 +50,7 @@ linters:
- unconvert - unconvert
- unparam - unparam
- wsl - wsl
- asasalint
#- errorlint causes stack overflow. TODO: recheck after each golangci update #- errorlint causes stack overflow. TODO: recheck after each golangci update
linters-settings: linters-settings:

View file

@ -75,7 +75,7 @@ lint:
lintci-deps: lintci-deps:
rm -f ./build/bin/golangci-lint rm -f ./build/bin/golangci-lint
curl -sSfL https://raw.githubusercontent.com/golangci/golangci-lint/master/install.sh | sh -s -- -b ./build/bin v1.46.0 curl -sSfL https://raw.githubusercontent.com/golangci/golangci-lint/master/install.sh | sh -s -- -b ./build/bin v1.48.0
goimports: goimports:
goimports -local "$(PACKAGE)" -w . goimports -local "$(PACKAGE)" -w .

View file

@ -21,7 +21,6 @@ import (
"crypto/ecdsa" "crypto/ecdsa"
"errors" "errors"
"io" "io"
"io/ioutil"
"math/big" "math/big"
"github.com/ethereum/go-ethereum/accounts" "github.com/ethereum/go-ethereum/accounts"
@ -45,14 +44,17 @@ var ErrNotAuthorized = errors.New("not authorized to sign this account")
// Deprecated: Use NewTransactorWithChainID instead. // Deprecated: Use NewTransactorWithChainID instead.
func NewTransactor(keyin io.Reader, passphrase string) (*TransactOpts, error) { func NewTransactor(keyin io.Reader, passphrase string) (*TransactOpts, error) {
log.Warn("WARNING: NewTransactor has been deprecated in favour of NewTransactorWithChainID") log.Warn("WARNING: NewTransactor has been deprecated in favour of NewTransactorWithChainID")
json, err := ioutil.ReadAll(keyin)
json, err := io.ReadAll(keyin)
if err != nil { if err != nil {
return nil, err return nil, err
} }
key, err := keystore.DecryptKey(json, passphrase) key, err := keystore.DecryptKey(json, passphrase)
if err != nil { if err != nil {
return nil, err return nil, err
} }
return NewKeyedTransactor(key.PrivateKey), nil return NewKeyedTransactor(key.PrivateKey), nil
} }
@ -106,7 +108,7 @@ func NewKeyedTransactor(key *ecdsa.PrivateKey) *TransactOpts {
// NewTransactorWithChainID is a utility method to easily create a transaction signer from // NewTransactorWithChainID is a utility method to easily create a transaction signer from
// an encrypted json key stream and the associated passphrase. // an encrypted json key stream and the associated passphrase.
func NewTransactorWithChainID(keyin io.Reader, passphrase string, chainID *big.Int) (*TransactOpts, error) { func NewTransactorWithChainID(keyin io.Reader, passphrase string, chainID *big.Int) (*TransactOpts, error) {
json, err := ioutil.ReadAll(keyin) json, err := io.ReadAll(keyin)
if err != nil { if err != nil {
return nil, err return nil, err
} }

View file

@ -18,7 +18,6 @@ package bind
import ( import (
"fmt" "fmt"
"io/ioutil"
"os" "os"
"os/exec" "os/exec"
"path/filepath" "path/filepath"
@ -1966,7 +1965,7 @@ func TestGolangBindings(t *testing.T) {
t.Skip("go sdk not found for testing") t.Skip("go sdk not found for testing")
} }
// Create a temporary workspace for the test suite // Create a temporary workspace for the test suite
ws, err := ioutil.TempDir("", "binding-test") ws, err := os.MkdirTemp("", "binding-test")
if err != nil { if err != nil {
t.Fatalf("failed to create temporary workspace: %v", err) t.Fatalf("failed to create temporary workspace: %v", err)
} }
@ -1990,7 +1989,7 @@ func TestGolangBindings(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("test %d: failed to generate binding: %v", i, err) t.Fatalf("test %d: failed to generate binding: %v", i, err)
} }
if err = ioutil.WriteFile(filepath.Join(pkg, strings.ToLower(tt.name)+".go"), []byte(bind), 0600); err != nil { if err = os.WriteFile(filepath.Join(pkg, strings.ToLower(tt.name)+".go"), []byte(bind), 0600); err != nil {
t.Fatalf("test %d: failed to write binding: %v", i, err) t.Fatalf("test %d: failed to write binding: %v", i, err)
} }
// Generate the test file with the injected test code // Generate the test file with the injected test code
@ -2006,7 +2005,7 @@ func TestGolangBindings(t *testing.T) {
%s %s
} }
`, tt.imports, tt.name, tt.tester) `, tt.imports, tt.name, tt.tester)
if err := ioutil.WriteFile(filepath.Join(pkg, strings.ToLower(tt.name)+"_test.go"), []byte(code), 0600); err != nil { if err := os.WriteFile(filepath.Join(pkg, strings.ToLower(tt.name)+"_test.go"), []byte(code), 0600); err != nil {
t.Fatalf("test %d: failed to write tests: %v", i, err) t.Fatalf("test %d: failed to write tests: %v", i, err)
} }
}) })

View file

@ -18,14 +18,12 @@ package main
import ( import (
"fmt" "fmt"
"io/ioutil"
"math/big" "math/big"
"os" "os"
"time" "time"
"gopkg.in/urfave/cli.v1"
"github.com/BurntSushi/toml" "github.com/BurntSushi/toml"
"gopkg.in/urfave/cli.v1"
"github.com/ethereum/go-ethereum/accounts/external" "github.com/ethereum/go-ethereum/accounts/external"
"github.com/ethereum/go-ethereum/accounts/keystore" "github.com/ethereum/go-ethereum/accounts/keystore"
@ -71,7 +69,7 @@ type gethConfig struct {
} }
func loadConfig(file string, cfg *gethConfig) error { func loadConfig(file string, cfg *gethConfig) error {
data, err := ioutil.ReadFile(file) data, err := os.ReadFile(file)
if err != nil { if err != nil {
return err return err
} }

View file

@ -12,6 +12,7 @@ import (
"sort" "sort"
"strconv" "strconv"
"sync" "sync"
"sync/atomic"
"time" "time"
lru "github.com/hashicorp/golang-lru" lru "github.com/hashicorp/golang-lru"
@ -97,6 +98,9 @@ var (
// errOutOfRangeChain is returned if an authorization list is attempted to // errOutOfRangeChain is returned if an authorization list is attempted to
// be modified via out-of-range or non-contiguous headers. // be modified via out-of-range or non-contiguous headers.
errOutOfRangeChain = errors.New("out of range or non-contiguous chain") errOutOfRangeChain = errors.New("out of range or non-contiguous chain")
errUncleDetected = errors.New("uncles not allowed")
errUnknownValidators = errors.New("unknown validators")
) )
// SignerFn is a signer callback function to request a header to be signed by a // SignerFn is a signer callback function to request a header to be signed by a
@ -210,9 +214,7 @@ type Bor struct {
recents *lru.ARCCache // Snapshots for recent block to speed up reorgs recents *lru.ARCCache // Snapshots for recent block to speed up reorgs
signatures *lru.ARCCache // Signatures of recent blocks to speed up mining signatures *lru.ARCCache // Signatures of recent blocks to speed up mining
signer common.Address // Ethereum address of the signing key authorizedSigner atomic.Pointer[signer] // Ethereum address and sign function of the signing key
signFn SignerFn // Signer function to authorize hashes with
lock sync.RWMutex // Protects the signer fields
ethAPI api.Caller ethAPI api.Caller
spanner Spanner spanner Spanner
@ -225,6 +227,11 @@ type Bor struct {
closeOnce sync.Once closeOnce sync.Once
} }
type signer struct {
signer common.Address // Ethereum address of the signing key
signFn SignerFn // Signer function to authorize hashes with
}
// New creates a Matic Bor consensus engine. // New creates a Matic Bor consensus engine.
func New( func New(
chainConfig *params.ChainConfig, chainConfig *params.ChainConfig,
@ -257,6 +264,14 @@ func New(
HeimdallClient: heimdallClient, HeimdallClient: heimdallClient,
} }
c.authorizedSigner.Store(&signer{
common.Address{},
func(_ accounts.Account, _ string, i []byte) ([]byte, error) {
// return an error to prevent panics
return nil, &UnauthorizedSignerError{0, common.Address{}.Bytes()}
},
})
// make sure we can decode all the GenesisAlloc in the BorConfig. // make sure we can decode all the GenesisAlloc in the BorConfig.
for key, genesisAlloc := range c.config.BlockAlloc { for key, genesisAlloc := range c.config.BlockAlloc {
if _, err := decodeGenesisAlloc(genesisAlloc); err != nil { if _, err := decodeGenesisAlloc(genesisAlloc); err != nil {
@ -572,7 +587,7 @@ func (c *Bor) snapshot(chain consensus.ChainHeaderReader, number uint64, hash co
// uncles as this consensus mechanism doesn't permit uncles. // uncles as this consensus mechanism doesn't permit uncles.
func (c *Bor) VerifyUncles(_ consensus.ChainReader, block *types.Block) error { func (c *Bor) VerifyUncles(_ consensus.ChainReader, block *types.Block) error {
if len(block.Uncles()) > 0 { if len(block.Uncles()) > 0 {
return errors.New("uncles not allowed") return errUncleDetected
} }
return nil return nil
@ -656,8 +671,10 @@ func (c *Bor) Prepare(chain consensus.ChainHeaderReader, header *types.Header) e
return err return err
} }
currentSigner := *c.authorizedSigner.Load()
// Set the correct difficulty // Set the correct difficulty
header.Difficulty = new(big.Int).SetUint64(Difficulty(snap.ValidatorSet, c.signer)) header.Difficulty = new(big.Int).SetUint64(Difficulty(snap.ValidatorSet, currentSigner.signer))
// Ensure the extra data has all it's components // Ensure the extra data has all it's components
if len(header.Extra) < extraVanity { if len(header.Extra) < extraVanity {
@ -670,7 +687,7 @@ func (c *Bor) Prepare(chain consensus.ChainHeaderReader, header *types.Header) e
if IsSprintStart(number+1, c.config.Sprint) { if IsSprintStart(number+1, c.config.Sprint) {
newValidators, err := c.spanner.GetCurrentValidators(context.Background(), header.ParentHash, number+1) newValidators, err := c.spanner.GetCurrentValidators(context.Background(), header.ParentHash, number+1)
if err != nil { if err != nil {
return errors.New("unknown validators") return errUnknownValidators
} }
// sort validator by address // sort validator by address
@ -695,8 +712,8 @@ func (c *Bor) Prepare(chain consensus.ChainHeaderReader, header *types.Header) e
var succession int var succession int
// if signer is not empty // if signer is not empty
if c.signer != (common.Address{}) { if currentSigner.signer != (common.Address{}) {
succession, err = snap.GetSignerSuccessionNumber(c.signer) succession, err = snap.GetSignerSuccessionNumber(currentSigner.signer)
if err != nil { if err != nil {
return err return err
} }
@ -774,7 +791,7 @@ func (c *Bor) changeContractCodeIfNeeded(headerNumber uint64, state *state.State
if blockNumber == strconv.FormatUint(headerNumber, 10) { if blockNumber == strconv.FormatUint(headerNumber, 10) {
allocs, err := decodeGenesisAlloc(genesisAlloc) allocs, err := decodeGenesisAlloc(genesisAlloc)
if err != nil { if err != nil {
return fmt.Errorf("failed to decode genesis alloc: %v", err) return fmt.Errorf("failed to decode genesis alloc: %w", err)
} }
for addr, account := range allocs { for addr, account := range allocs {
@ -838,12 +855,11 @@ func (c *Bor) FinalizeAndAssemble(chain consensus.ChainHeaderReader, header *typ
// Authorize injects a private key into the consensus engine to mint new blocks // Authorize injects a private key into the consensus engine to mint new blocks
// with. // with.
func (c *Bor) Authorize(signer common.Address, signFn SignerFn) { func (c *Bor) Authorize(currentSigner common.Address, signFn SignerFn) {
c.lock.Lock() c.authorizedSigner.Store(&signer{
defer c.lock.Unlock() signer: currentSigner,
signFn: signFn,
c.signer = signer })
c.signFn = signFn
} }
// Seal implements consensus.Engine, attempting to create a sealed block using // Seal implements consensus.Engine, attempting to create a sealed block using
@ -860,10 +876,9 @@ func (c *Bor) Seal(chain consensus.ChainHeaderReader, block *types.Block, result
log.Info("Sealing paused, waiting for transactions") log.Info("Sealing paused, waiting for transactions")
return nil return nil
} }
// Don't hold the signer fields for the entire sealing procedure // Don't hold the signer fields for the entire sealing procedure
c.lock.RLock() currentSigner := *c.authorizedSigner.Load()
signer, signFn := c.signer, c.signFn
c.lock.RUnlock()
snap, err := c.snapshot(chain, number-1, header.ParentHash, nil) snap, err := c.snapshot(chain, number-1, header.ParentHash, nil)
if err != nil { if err != nil {
@ -871,12 +886,12 @@ func (c *Bor) Seal(chain consensus.ChainHeaderReader, block *types.Block, result
} }
// Bail out if we're unauthorized to sign a block // Bail out if we're unauthorized to sign a block
if !snap.ValidatorSet.HasAddress(signer) { if !snap.ValidatorSet.HasAddress(currentSigner.signer) {
// Check the UnauthorizedSignerError.Error() msg to see why we pass number-1 // Check the UnauthorizedSignerError.Error() msg to see why we pass number-1
return &UnauthorizedSignerError{number - 1, signer.Bytes()} return &UnauthorizedSignerError{number - 1, currentSigner.signer.Bytes()}
} }
successionNumber, err := snap.GetSignerSuccessionNumber(signer) successionNumber, err := snap.GetSignerSuccessionNumber(currentSigner.signer)
if err != nil { if err != nil {
return err return err
} }
@ -887,7 +902,7 @@ func (c *Bor) Seal(chain consensus.ChainHeaderReader, block *types.Block, result
wiggle := time.Duration(successionNumber) * time.Duration(c.config.CalculateBackupMultiplier(number)) * time.Second wiggle := time.Duration(successionNumber) * time.Duration(c.config.CalculateBackupMultiplier(number)) * time.Second
// Sign all the things! // Sign all the things!
err = Sign(signFn, signer, header, c.config) err = Sign(currentSigner.signFn, currentSigner.signer, header, c.config)
if err != nil { if err != nil {
return err return err
} }
@ -949,7 +964,7 @@ func (c *Bor) CalcDifficulty(chain consensus.ChainHeaderReader, _ uint64, parent
return nil return nil
} }
return new(big.Int).SetUint64(Difficulty(snap.ValidatorSet, c.signer)) return new(big.Int).SetUint64(Difficulty(snap.ValidatorSet, c.authorizedSigner.Load().signer))
} }
// SealHash returns the hash of a block prior to it being sealed. // SealHash returns the hash of a block prior to it being sealed.

View file

@ -3,10 +3,10 @@ package heimdallgrpc
import ( import (
"context" "context"
proto "github.com/maticnetwork/polyproto/heimdall"
"github.com/ethereum/go-ethereum/common" "github.com/ethereum/go-ethereum/common"
"github.com/ethereum/go-ethereum/consensus/bor/clerk" "github.com/ethereum/go-ethereum/consensus/bor/clerk"
proto "github.com/maticnetwork/polyproto/heimdall"
) )
func (h *HeimdallGRPCClient) StateSyncEvents(ctx context.Context, fromID uint64, to int64) ([]*clerk.EventRecordWithTime, error) { func (h *HeimdallGRPCClient) StateSyncEvents(ctx context.Context, fromID uint64, to int64) ([]*clerk.EventRecordWithTime, error) {
@ -18,13 +18,19 @@ func (h *HeimdallGRPCClient) StateSyncEvents(ctx context.Context, fromID uint64,
Limit: uint64(stateFetchLimit), Limit: uint64(stateFetchLimit),
} }
res, err := h.client.StateSyncEvents(ctx, req) var (
res proto.Heimdall_StateSyncEventsClient
events *proto.StateSyncEventsResponse
err error
)
res, err = h.client.StateSyncEvents(ctx, req)
if err != nil { if err != nil {
return nil, err return nil, err
} }
for { for {
events, err := res.Recv() events, err = res.Recv()
if err != nil { if err != nil {
break break
} }
@ -45,5 +51,5 @@ func (h *HeimdallGRPCClient) StateSyncEvents(ctx context.Context, fromID uint64,
} }
} }
return eventRecords, nil return eventRecords, err
} }

View file

@ -207,7 +207,7 @@ type BlockChain interface {
} }
// New creates a new downloader to fetch hashes and blocks from remote peers. // New creates a new downloader to fetch hashes and blocks from remote peers.
//nolint: staticcheck // nolint: staticcheck
func New(checkpoint uint64, stateDb ethdb.Database, mux *event.TypeMux, chain BlockChain, lightchain LightChain, dropPeer peerDropFn, success func(), whitelistService ethereum.ChainValidator) *Downloader { func New(checkpoint uint64, stateDb ethdb.Database, mux *event.TypeMux, chain BlockChain, lightchain LightChain, dropPeer peerDropFn, success func(), whitelistService ethereum.ChainValidator) *Downloader {
if lightchain == nil { if lightchain == nil {
lightchain = chain lightchain = chain
@ -729,9 +729,11 @@ func (d *Downloader) fetchHead(p *peerConnection) (head *types.Header, pivot *ty
// calculateRequestSpan calculates what headers to request from a peer when trying to determine the // calculateRequestSpan calculates what headers to request from a peer when trying to determine the
// common ancestor. // common ancestor.
// It returns parameters to be used for peer.RequestHeadersByNumber: // It returns parameters to be used for peer.RequestHeadersByNumber:
// from - starting block number //
// count - number of headers to request // from - starting block number
// skip - number of headers to skip // count - number of headers to request
// skip - number of headers to skip
//
// and also returns 'max', the last block which is expected to be returned by the remote peers, // and also returns 'max', the last block which is expected to be returned by the remote peers,
// given the (from,count,skip) // given the (from,count,skip)
func calculateRequestSpan(remoteHeight, localHeight uint64) (int64, int, int, uint64) { func calculateRequestSpan(remoteHeight, localHeight uint64) (int64, int, int, uint64) {

View file

@ -48,7 +48,11 @@ func (b *TestBackend) GetBorBlockReceipt(ctx context.Context, hash common.Hash)
func (b *TestBackend) GetBorBlockLogs(ctx context.Context, hash common.Hash) ([]*types.Log, error) { func (b *TestBackend) GetBorBlockLogs(ctx context.Context, hash common.Hash) ([]*types.Log, error) {
receipt, err := b.GetBorBlockReceipt(ctx, hash) receipt, err := b.GetBorBlockReceipt(ctx, hash)
if receipt == nil || err != nil { if err != nil {
return []*types.Log{}, err
}
if receipt == nil {
return []*types.Log{}, nil return []*types.Log{}, nil
} }

2
go.mod
View file

@ -1,6 +1,6 @@
module github.com/ethereum/go-ethereum module github.com/ethereum/go-ethereum
go 1.18 go 1.19
require ( require (
github.com/Azure/azure-sdk-for-go/sdk/storage/azblob v0.3.0 github.com/Azure/azure-sdk-for-go/sdk/storage/azblob v0.3.0

View file

@ -2,13 +2,13 @@ package server
import ( import (
"fmt" "fmt"
"io/ioutil" "os"
"github.com/BurntSushi/toml" "github.com/BurntSushi/toml"
) )
func readLegacyConfig(path string) (*Config, error) { func readLegacyConfig(path string) (*Config, error) {
data, err := ioutil.ReadFile(path) data, err := os.ReadFile(path)
tomlData := string(data) tomlData := string(data)
if err != nil { if err != nil {

View file

@ -26,12 +26,13 @@ import (
// Handler returns a log handler which logs to the unit test log of t. // Handler returns a log handler which logs to the unit test log of t.
func Handler(t *testing.T, level log.Lvl) log.Handler { func Handler(t *testing.T, level log.Lvl) log.Handler {
return log.LvlFilterHandler(level, &handler{t, log.TerminalFormat(false)}) return log.LvlFilterHandler(level, &handler{t, log.TerminalFormat(false), level})
} }
type handler struct { type handler struct {
t *testing.T t *testing.T
fmt log.Format fmt log.Format
lvl log.Lvl
} }
func (h *handler) Log(r *log.Record) error { func (h *handler) Log(r *log.Record) error {
@ -39,6 +40,10 @@ func (h *handler) Log(r *log.Record) error {
return nil return nil
} }
func (h *handler) Level() log.Lvl {
return h.lvl
}
// logger implements log.Logger such that all output goes to the unit test log via // logger implements log.Logger such that all output goes to the unit test log via
// t.Logf(). All methods in between logger.Trace, logger.Debug, etc. are marked as test // t.Logf(). All methods in between logger.Trace, logger.Debug, etc. are marked as test
// helpers, so the file and line number in unit test output correspond to the call site // helpers, so the file and line number in unit test output correspond to the call site
@ -59,6 +64,9 @@ func (h *bufHandler) Log(r *log.Record) error {
h.buf = append(h.buf, r) h.buf = append(h.buf, r)
return nil return nil
} }
func (h *bufHandler) Level() log.Lvl {
return log.LvlTrace
}
// Logger returns a logger which logs to the unit test log of t. // Logger returns a logger which logs to the unit test log of t.
func Logger(t *testing.T, level log.Lvl) log.Logger { func Logger(t *testing.T, level log.Lvl) log.Logger {

View file

@ -17,18 +17,26 @@ import (
// them to achieve the logging structure that suits your applications. // them to achieve the logging structure that suits your applications.
type Handler interface { type Handler interface {
Log(r *Record) error Log(r *Record) error
Level() Lvl
} }
// FuncHandler returns a Handler that logs records with the given // FuncHandler returns a Handler that logs records with the given
// function. // function.
func FuncHandler(fn func(r *Record) error) Handler { func FuncHandler(fn func(r *Record) error, lvl Lvl) Handler {
return funcHandler(fn) return funcHandler{fn, lvl}
} }
type funcHandler func(r *Record) error type funcHandler struct {
log func(r *Record) error
lvl Lvl
}
func (h funcHandler) Log(r *Record) error { func (h funcHandler) Log(r *Record) error {
return h(r) return h.log(r)
}
func (h funcHandler) Level() Lvl {
return h.lvl
} }
// StreamHandler writes log records to an io.Writer // StreamHandler writes log records to an io.Writer
@ -42,7 +50,7 @@ func StreamHandler(wr io.Writer, fmtr Format) Handler {
h := FuncHandler(func(r *Record) error { h := FuncHandler(func(r *Record) error {
_, err := wr.Write(fmtr.Format(r)) _, err := wr.Write(fmtr.Format(r))
return err return err
}) }, LvlTrace)
return LazyHandler(SyncHandler(h)) return LazyHandler(SyncHandler(h))
} }
@ -55,7 +63,7 @@ func SyncHandler(h Handler) Handler {
defer mu.Unlock() defer mu.Unlock()
mu.Lock() mu.Lock()
return h.Log(r) return h.Log(r)
}) }, h.Level())
} }
// FileHandler returns a handler which writes log records to the give file // FileHandler returns a handler which writes log records to the give file
@ -99,7 +107,7 @@ func CallerFileHandler(h Handler) Handler {
return FuncHandler(func(r *Record) error { return FuncHandler(func(r *Record) error {
r.Ctx = append(r.Ctx, "caller", fmt.Sprint(r.Call)) r.Ctx = append(r.Ctx, "caller", fmt.Sprint(r.Call))
return h.Log(r) return h.Log(r)
}) }, h.Level())
} }
// CallerFuncHandler returns a Handler that adds the calling function name to // CallerFuncHandler returns a Handler that adds the calling function name to
@ -108,7 +116,7 @@ func CallerFuncHandler(h Handler) Handler {
return FuncHandler(func(r *Record) error { return FuncHandler(func(r *Record) error {
r.Ctx = append(r.Ctx, "fn", formatCall("%+n", r.Call)) r.Ctx = append(r.Ctx, "fn", formatCall("%+n", r.Call))
return h.Log(r) return h.Log(r)
}) }, h.Level())
} }
// This function is here to please go vet on Go < 1.8. // This function is here to please go vet on Go < 1.8.
@ -128,29 +136,28 @@ func CallerStackHandler(format string, h Handler) Handler {
r.Ctx = append(r.Ctx, "stack", fmt.Sprintf(format, s)) r.Ctx = append(r.Ctx, "stack", fmt.Sprintf(format, s))
} }
return h.Log(r) return h.Log(r)
}) }, h.Level())
} }
// FilterHandler returns a Handler that only writes records to the // FilterHandler returns a Handler that only writes records to the
// wrapped Handler if the given function evaluates true. For example, // wrapped Handler if the given function evaluates true. For example,
// to only log records where the 'err' key is not nil: // to only log records where the 'err' key is not nil:
// //
// logger.SetHandler(FilterHandler(func(r *Record) bool { // logger.SetHandler(FilterHandler(func(r *Record) bool {
// for i := 0; i < len(r.Ctx); i += 2 { // for i := 0; i < len(r.Ctx); i += 2 {
// if r.Ctx[i] == "err" { // if r.Ctx[i] == "err" {
// return r.Ctx[i+1] != nil // return r.Ctx[i+1] != nil
// } // }
// } // }
// return false // return false
// }, h)) // }, h))
//
func FilterHandler(fn func(r *Record) bool, h Handler) Handler { func FilterHandler(fn func(r *Record) bool, h Handler) Handler {
return FuncHandler(func(r *Record) error { return FuncHandler(func(r *Record) error {
if fn(r) { if fn(r) {
return h.Log(r) return h.Log(r)
} }
return nil return nil
}) }, h.Level())
} }
// MatchFilterHandler returns a Handler that only writes records // MatchFilterHandler returns a Handler that only writes records
@ -158,8 +165,7 @@ func FilterHandler(fn func(r *Record) bool, h Handler) Handler {
// context matches the value. For example, to only log records // context matches the value. For example, to only log records
// from your ui package: // from your ui package:
// //
// log.MatchFilterHandler("pkg", "app/ui", log.StdoutHandler) // log.MatchFilterHandler("pkg", "app/ui", log.StdoutHandler)
//
func MatchFilterHandler(key string, value interface{}, h Handler) Handler { func MatchFilterHandler(key string, value interface{}, h Handler) Handler {
return FilterHandler(func(r *Record) (pass bool) { return FilterHandler(func(r *Record) (pass bool) {
switch key { switch key {
@ -185,8 +191,7 @@ func MatchFilterHandler(key string, value interface{}, h Handler) Handler {
// level to the wrapped Handler. For example, to only // level to the wrapped Handler. For example, to only
// log Error/Crit records: // log Error/Crit records:
// //
// log.LvlFilterHandler(log.LvlError, log.StdoutHandler) // log.LvlFilterHandler(log.LvlError, log.StdoutHandler)
//
func LvlFilterHandler(maxLvl Lvl, h Handler) Handler { func LvlFilterHandler(maxLvl Lvl, h Handler) Handler {
return FilterHandler(func(r *Record) (pass bool) { return FilterHandler(func(r *Record) (pass bool) {
return r.Lvl <= maxLvl return r.Lvl <= maxLvl
@ -198,10 +203,9 @@ func LvlFilterHandler(maxLvl Lvl, h Handler) Handler {
// to different locations. For example, to log to a file and // to different locations. For example, to log to a file and
// standard error: // standard error:
// //
// log.MultiHandler( // log.MultiHandler(
// log.Must.FileHandler("/var/log/app.log", log.LogfmtFormat()), // log.Must.FileHandler("/var/log/app.log", log.LogfmtFormat()),
// log.StderrHandler) // log.StderrHandler)
//
func MultiHandler(hs ...Handler) Handler { func MultiHandler(hs ...Handler) Handler {
return FuncHandler(func(r *Record) error { return FuncHandler(func(r *Record) error {
for _, h := range hs { for _, h := range hs {
@ -209,7 +213,7 @@ func MultiHandler(hs ...Handler) Handler {
h.Log(r) h.Log(r)
} }
return nil return nil
}) }, LvlDebug)
} }
// FailoverHandler writes all log records to the first handler // FailoverHandler writes all log records to the first handler
@ -219,10 +223,10 @@ func MultiHandler(hs ...Handler) Handler {
// to writing to a file if the network fails, and then to // to writing to a file if the network fails, and then to
// standard out if the file write fails: // standard out if the file write fails:
// //
// log.FailoverHandler( // log.FailoverHandler(
// log.Must.NetHandler("tcp", ":9090", log.JSONFormat()), // log.Must.NetHandler("tcp", ":9090", log.JSONFormat()),
// log.Must.FileHandler("/var/log/app.log", log.LogfmtFormat()), // log.Must.FileHandler("/var/log/app.log", log.LogfmtFormat()),
// log.StdoutHandler) // log.StdoutHandler)
// //
// All writes that do not go to the first handler will add context with keys of // All writes that do not go to the first handler will add context with keys of
// the form "failover_err_{idx}" which explain the error encountered while // the form "failover_err_{idx}" which explain the error encountered while
@ -239,17 +243,17 @@ func FailoverHandler(hs ...Handler) Handler {
} }
return err return err
}) }, LvlTrace)
} }
// ChannelHandler writes all records to the given channel. // ChannelHandler writes all records to the given channel.
// It blocks if the channel is full. Useful for async processing // It blocks if the channel is full. Useful for async processing
// of log messages, it's used by BufferedHandler. // of log messages, it's used by BufferedHandler.
func ChannelHandler(recs chan<- *Record) Handler { func ChannelHandler(recs chan<- *Record, lvl Lvl) Handler {
return FuncHandler(func(r *Record) error { return FuncHandler(func(r *Record) error {
recs <- r recs <- r
return nil return nil
}) }, lvl)
} }
// BufferedHandler writes all records to a buffered // BufferedHandler writes all records to a buffered
@ -264,7 +268,8 @@ func BufferedHandler(bufSize int, h Handler) Handler {
_ = h.Log(m) _ = h.Log(m)
} }
}() }()
return ChannelHandler(recs)
return ChannelHandler(recs, h.Level())
} }
// LazyHandler writes all values to the wrapped handler after evaluating // LazyHandler writes all values to the wrapped handler after evaluating
@ -297,7 +302,7 @@ func LazyHandler(h Handler) Handler {
} }
return h.Log(r) return h.Log(r)
}) }, h.Level())
} }
func evaluateLazy(lz Lazy) (interface{}, error) { func evaluateLazy(lz Lazy) (interface{}, error) {
@ -333,7 +338,7 @@ func evaluateLazy(lz Lazy) (interface{}, error) {
func DiscardHandler() Handler { func DiscardHandler() Handler {
return FuncHandler(func(r *Record) error { return FuncHandler(func(r *Record) error {
return nil return nil
}) }, LvlDiscard)
} }
// Must provides the following Handler creation functions // Must provides the following Handler creation functions

View file

@ -82,14 +82,14 @@ func (h *GlogHandler) Verbosity(level Lvl) {
// //
// For instance: // For instance:
// //
// pattern="gopher.go=3" // pattern="gopher.go=3"
// sets the V level to 3 in all Go files named "gopher.go" // sets the V level to 3 in all Go files named "gopher.go"
// //
// pattern="foo=3" // pattern="foo=3"
// sets V to 3 in all files of any packages whose import path ends in "foo" // sets V to 3 in all files of any packages whose import path ends in "foo"
// //
// pattern="foo/*=3" // pattern="foo/*=3"
// sets V to 3 in all files of any packages whose import path contains "foo" // sets V to 3 in all files of any packages whose import path contains "foo"
func (h *GlogHandler) Vmodule(ruleset string) error { func (h *GlogHandler) Vmodule(ruleset string) error {
var filter []pattern var filter []pattern
for _, rule := range strings.Split(ruleset, ",") { for _, rule := range strings.Split(ruleset, ",") {
@ -230,3 +230,7 @@ func (h *GlogHandler) Log(r *Record) error {
} }
return nil return nil
} }
func (h *GlogHandler) Level() Lvl {
return Lvl(atomic.LoadUint32(&h.level))
}

27
log/handler_go119.go Normal file
View file

@ -0,0 +1,27 @@
//+go:build go1.19
package log
import "sync/atomic"
// swapHandler wraps another handler that may be swapped out
// dynamically at runtime in a thread-safe fashion.
type swapHandler struct {
handler atomic.Pointer[Handler]
}
func (h *swapHandler) Log(r *Record) error {
return (*h.handler.Load()).Log(r)
}
func (h *swapHandler) Swap(newHandler Handler) {
h.handler.Store(&newHandler)
}
func (h *swapHandler) Get() Handler {
return *h.handler.Load()
}
func (h *swapHandler) Level() Lvl {
return (*h.handler.Load()).Level()
}

View file

@ -1,5 +1,4 @@
//go:build go1.4 //go:build !go1.19
// +build go1.4
package log package log

View file

@ -18,7 +18,8 @@ const skipLevel = 2
type Lvl int type Lvl int
const ( const (
LvlCrit Lvl = iota LvlDiscard Lvl = -1
LvlCrit Lvl = iota
LvlError LvlError
LvlWarn LvlWarn
LvlInfo LvlInfo
@ -131,6 +132,10 @@ type logger struct {
} }
func (l *logger) write(msg string, lvl Lvl, ctx []interface{}, skip int) { func (l *logger) write(msg string, lvl Lvl, ctx []interface{}, skip int) {
if l.h.Level() < lvl {
return
}
l.h.Log(&Record{ l.h.Log(&Record{
Time: time.Now(), Time: time.Now(),
Lvl: lvl, Lvl: lvl,

View file

@ -45,7 +45,7 @@ func sharedSyslog(fmtr Format, sysWr *syslog.Writer, err error) (Handler, error)
s := strings.TrimSpace(string(fmtr.Format(r))) s := strings.TrimSpace(string(fmtr.Format(r)))
return syslogFn(s) return syslogFn(s)
}) }, LvlTrace)
return LazyHandler(&closingHandler{sysWr, h}), nil return LazyHandler(&closingHandler{sysWr, h}), nil
} }

View file

@ -562,7 +562,7 @@ func startLocalhostV4(t *testing.T, cfg Config) *UDPv4 {
cfg.Log.SetHandler(log.FuncHandler(func(r *log.Record) error { cfg.Log.SetHandler(log.FuncHandler(func(r *log.Record) error {
t.Logf("%s %s", lprefix, lfmt.Format(r)) t.Logf("%s %s", lprefix, lfmt.Format(r))
return nil return nil
})) }, log.LvlTrace))
// Listen. // Listen.
socket, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IP{127, 0, 0, 1}}) socket, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IP{127, 0, 0, 1}})

View file

@ -83,7 +83,7 @@ func startLocalhostV5(t *testing.T, cfg Config) *UDPv5 {
cfg.Log.SetHandler(log.FuncHandler(func(r *log.Record) error { cfg.Log.SetHandler(log.FuncHandler(func(r *log.Record) error {
t.Logf("%s %s", lprefix, lfmt.Format(r)) t.Logf("%s %s", lprefix, lfmt.Format(r))
return nil return nil
})) }, log.LvlTrace))
// Listen. // Listen.
socket, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IP{127, 0, 0, 1}}) socket, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IP{127, 0, 0, 1}})