Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 55 additions & 2 deletions common/body_storage.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,15 @@ type BodyStorage interface {
Size() int64
// IsDisk 是否是磁盘存储
IsDisk() bool
// NewReader returns an independent reader positioned at the start.
NewReader() (io.ReadCloser, error)
}

// ReplayableBody exposes body metadata without transferring storage ownership.
type ReplayableBody interface {
io.Reader
Size() int64
NewReader() (io.ReadCloser, error)
}

// ErrStorageClosed 存储已关闭错误
Expand Down Expand Up @@ -80,6 +89,15 @@ func (m *memoryStorage) Bytes() ([]byte, error) {
return m.data, nil
}

func (m *memoryStorage) NewReader() (io.ReadCloser, error) {
m.mu.Lock()
defer m.mu.Unlock()
if atomic.LoadInt32(&m.closed) == 1 {
return nil, ErrStorageClosed
}
return io.NopCloser(bytes.NewReader(m.data)), nil
}

func (m *memoryStorage) Size() int64 {
return m.size
}
Expand Down Expand Up @@ -229,6 +247,19 @@ func (d *diskStorage) Bytes() ([]byte, error) {
return data, nil
}

func (d *diskStorage) NewReader() (io.ReadCloser, error) {
d.mu.Lock()
defer d.mu.Unlock()
if atomic.LoadInt32(&d.closed) == 1 {
return nil, ErrStorageClosed
}
file, err := os.Open(d.filePath)
if err != nil {
return nil, fmt.Errorf("failed to open body cache file for replay: %w", err)
}
return file, nil
}

func (d *diskStorage) Size() int64 {
return d.size
}
Expand Down Expand Up @@ -302,9 +333,31 @@ func CreateBodyStorageFromReader(reader io.Reader, contentLength int64, maxBytes
return storage, nil
}

// ReaderOnly wraps an io.Reader to hide io.Closer, preventing http.NewRequest
// from type-asserting io.ReadCloser and closing the underlying BodyStorage.
type replayableBodyReader struct {
storage BodyStorage
}

func (r replayableBodyReader) Read(p []byte) (int, error) {
return r.storage.Read(p)
}

func (r replayableBodyReader) Size() int64 {
return r.storage.Size()
}

func (r replayableBodyReader) NewReader() (io.ReadCloser, error) {
return r.storage.NewReader()
}

func NewReplayableBodyReader(storage BodyStorage) ReplayableBody {
return replayableBodyReader{storage: storage}
}

// ReaderOnly hides io.Closer while preserving replay metadata for BodyStorage.
func ReaderOnly(r io.Reader) io.Reader {
if storage, ok := r.(BodyStorage); ok {
return NewReplayableBodyReader(storage)
}
return struct{ io.Reader }{r}
}

Expand Down
53 changes: 53 additions & 0 deletions common/body_storage_replay_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
package common

import (
"io"
"testing"

"github.com/stretchr/testify/require"
)

func requireIndependentBodyReaders(t *testing.T, storage BodyStorage, payload []byte) {
t.Helper()
first, err := storage.NewReader()
require.NoError(t, err)
defer first.Close()
second, err := storage.NewReader()
require.NoError(t, err)
defer second.Close()

prefix := make([]byte, 4)
_, err = io.ReadFull(first, prefix)
require.NoError(t, err)
secondBody, err := io.ReadAll(second)
require.NoError(t, err)
require.Equal(t, payload, secondBody)
firstRest, err := io.ReadAll(first)
require.NoError(t, err)
require.Equal(t, payload[4:], firstRest)
}

func TestBodyStorageNewReaderUsesIndependentCursors(t *testing.T) {
payload := []byte("independent replay payload")

t.Run("memory", func(t *testing.T) {
storage := newMemoryStorage(payload)
defer storage.Close()
requireIndependentBodyReaders(t, storage, payload)
})

t.Run("disk", func(t *testing.T) {
storage, err := newDiskStorage(payload, "")
require.NoError(t, err)
defer storage.Close()
requireIndependentBodyReaders(t, storage, payload)
})
}

func TestBodyStorageNewReaderRejectsClosedStorage(t *testing.T) {
storage := newMemoryStorage([]byte("closed"))
require.NoError(t, storage.Close())
reader, err := storage.NewReader()
require.Nil(t, reader)
require.ErrorIs(t, err, ErrStorageClosed)
}
27 changes: 24 additions & 3 deletions common/email.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,21 +5,27 @@ import (
"encoding/base64"
"errors"
"fmt"
"net"
"net/smtp"
"slices"
"strings"
"time"
)

var ErrSMTPStartTLSUnsupported = errors.New("smtp server does not support STARTTLS")
var smtpOperationTimeout = 30 * time.Second

func generateMessageID() (string, error) {
split := strings.Split(SMTPFrom, "@")
if len(split) < 2 {
return "", fmt.Errorf("invalid SMTP account")
}
domain := strings.Split(SMTPFrom, "@")[1]
return fmt.Sprintf("<%d.%s@%s>", time.Now().UnixNano(), GetRandomString(12), domain), nil
randomPart, err := GenerateRandomCharsKey(12)
if err != nil {
return "", fmt.Errorf("generate message ID: %w", err)
}
return fmt.Sprintf("<%d.%s@%s>", time.Now().UnixNano(), randomPart, domain), nil
}

func shouldUseSMTPLoginAuth() bool {
Expand Down Expand Up @@ -59,11 +65,17 @@ func smtpTLSConfig() *tls.Config {
}

func newSMTPClient(addr string) (*smtp.Client, error) {
dialer := &net.Dialer{Timeout: smtpOperationTimeout}
deadline := time.Now().Add(smtpOperationTimeout)
if SMTPSSLEnabled || SMTPPort == 465 {
conn, err := tls.Dial("tcp", addr, smtpTLSConfig())
conn, err := tls.DialWithDialer(dialer, "tcp", addr, smtpTLSConfig())
if err != nil {
return nil, err
}
if err := conn.SetDeadline(deadline); err != nil {
_ = conn.Close()
return nil, err
}
client, err := smtp.NewClient(conn, SMTPServer)
if err != nil {
_ = conn.Close()
Expand All @@ -72,8 +84,17 @@ func newSMTPClient(addr string) (*smtp.Client, error) {
return client, nil
}

client, err := smtp.Dial(addr)
conn, err := dialer.Dial("tcp", addr)
if err != nil {
return nil, err
}
if err := conn.SetDeadline(deadline); err != nil {
_ = conn.Close()
return nil, err
}
client, err := smtp.NewClient(conn, SMTPServer)
if err != nil {
_ = conn.Close()
return nil, err
}

Expand Down
47 changes: 47 additions & 0 deletions common/email_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -220,6 +220,53 @@ func TestSendEmailSkipsAuthWhenServerDoesNotAdvertiseAuth(t *testing.T) {
}
}

func TestNewSMTPClientUsesOperationDeadline(t *testing.T) {
restore := preserveSMTPGlobals()
defer restore()

listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer listener.Close()

SMTPServer = "localhost"
SMTPPort = mustPort(t, listener.Addr().String())
SMTPSSLEnabled = false
SMTPStartTLSEnabled = false

oldTimeout := smtpOperationTimeout
smtpOperationTimeout = 50 * time.Millisecond
defer func() { smtpOperationTimeout = oldTimeout }()

done := make(chan error, 1)
go func() {
conn, acceptErr := listener.Accept()
if acceptErr != nil {
done <- acceptErr
return
}
defer conn.Close()
time.Sleep(200 * time.Millisecond)
done <- nil
}()

started := time.Now()
client, err := newSMTPClient(listener.Addr().String())
if client != nil {
_ = client.Close()
}
if err == nil {
t.Fatal("expected SMTP greeting timeout")
}
if time.Since(started) >= time.Second {
t.Fatalf("SMTP timeout took too long: %s", time.Since(started))
}
if err := waitSMTPServer(done); err != nil {
t.Fatal(err)
}
}

func preserveSMTPGlobals() func() {
old := struct {
server string
Expand Down
9 changes: 5 additions & 4 deletions common/go-channel.go
Original file line number Diff line number Diff line change
Expand Up @@ -34,12 +34,13 @@ func SafeSendString(ch chan string, value string) (closed bool) {
return false
}

// SafeSendStringTimeout send, return true, else return false
func SafeSendStringTimeout(ch chan string, value string, timeout int) (closed bool) {
// SafeSendStringTimeout returns true only when the value is sent before the timeout.
// A timeout or a closed channel returns false.
func SafeSendStringTimeout(ch chan string, value string, timeout int) (sent bool) {
defer func() {
// Recover from panic if one occured. A panic would mean the channel was closed.
// Recover from panic if the channel was closed between selection and send.
if recover() != nil {
closed = false
sent = false
}
}()

Expand Down
29 changes: 29 additions & 0 deletions common/go_channel_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
package common

import (
"testing"
"time"

"github.com/stretchr/testify/require"
)

func TestSafeSendStringTimeoutReportsSendOutcome(t *testing.T) {
t.Run("sent", func(t *testing.T) {
ch := make(chan string, 1)
require.True(t, SafeSendStringTimeout(ch, "value", 1))
require.Equal(t, "value", <-ch)
})

t.Run("timeout", func(t *testing.T) {
ch := make(chan string)
started := time.Now()
require.False(t, SafeSendStringTimeout(ch, "value", 0))
require.Less(t, time.Since(started), time.Second)
})

t.Run("closed", func(t *testing.T) {
ch := make(chan string)
close(ch)
require.False(t, SafeSendStringTimeout(ch, "value", 1))
})
}
13 changes: 4 additions & 9 deletions common/init.go
Original file line number Diff line number Diff line change
Expand Up @@ -93,15 +93,10 @@ func InitEnv() {
SessionCookieSecure = GetEnvOrDefaultBool("SESSION_COOKIE_SECURE", false)
SMTPStartTLSEnabled = GetEnvOrDefaultBool("SMTP_STARTTLS_ENABLE", GetEnvOrDefaultBool("SMTP_STARTTLS_ENABLED", false))
SMTPInsecureSkipVerify = GetEnvOrDefaultBool("SMTP_INSECURE_SKIP_VERIFY", GetEnvOrDefaultBool("SMTP_TLS_INSECURE_SKIP_VERIFY", false))
TrustedProxies = append([]string(nil), DefaultTrustedProxies...)
if trustedProxies := strings.TrimSpace(os.Getenv("TRUSTED_PROXIES")); trustedProxies != "" {
TrustedProxies = nil
for _, proxy := range strings.Split(trustedProxies, ",") {
proxy = strings.TrimSpace(proxy)
if proxy != "" {
TrustedProxies = append(TrustedProxies, proxy)
}
}
var trustedProxyErr error
TrustedProxies, trustedProxyErr = parseTrustedProxies(os.Getenv("TRUSTED_PROXIES"))
if trustedProxyErr != nil {
log.Fatal(trustedProxyErr)
}
IsMasterNode = os.Getenv("NODE_TYPE") != "slave"
initNodeNameIdentity()
Expand Down
Loading
Loading