package mail import ( "bufio" "context" "crypto/ecdsa" "crypto/elliptic" "crypto/rand" "crypto/tls" "crypto/x509" "encoding/base64" "errors" "fmt" "io" "math/big" "net" "net/http" "strings" "testing" "time" "github.com/emersion/go-imap/v2/imapclient" "golang.org/x/net/proxy" "thuanle.me/claw-email-bridge/internal/config" ) func TestConnectAndWatch_UsesDirectDial_WhenProxyUnset(t *testing.T) { watcher := &IMAPWatcher{ cfg: &config.Config{ IMAPHost: "imap.example.com", IMAPPort: "993", IMAPUser: "u", IMAPPass: "p", }, } var directCalled bool watcher.dialIMAP = func(addr string, _ *imapclient.Options) (*imapclient.Client, error) { directCalled = true return nil, errors.New("stop") } watcher.dialIMAPViaProxy = func(addr, proxyURL string, _ *imapclient.Options) (*imapclient.Client, error) { t.Fatalf("did not expect proxy dial, got addr=%s proxy=%s", addr, proxyURL) return nil, nil } err := watcher.connectAndWatch(context.Background()) if err == nil || !strings.Contains(err.Error(), "stop") { t.Fatalf("expected stop error, got %v", err) } if !directCalled { t.Fatal("expected direct dial path") } } func TestConnectAndWatch_UsesProxyDial_WhenProxySet(t *testing.T) { watcher := &IMAPWatcher{ cfg: &config.Config{ IMAPHost: "imap.example.com", IMAPPort: "993", IMAPUser: "u", IMAPPass: "p", IMAPProxyURL: "socks5://127.0.0.1:1080", }, } var proxyCalled bool watcher.dialIMAP = func(addr string, _ *imapclient.Options) (*imapclient.Client, error) { t.Fatalf("did not expect direct dial, got addr=%s", addr) return nil, nil } watcher.dialIMAPViaProxy = func(addr, proxyURL string, _ *imapclient.Options) (*imapclient.Client, error) { proxyCalled = true return nil, errors.New("stop") } err := watcher.connectAndWatch(context.Background()) if err == nil || !strings.Contains(err.Error(), "stop") { t.Fatalf("expected stop error, got %v", err) } if !proxyCalled { t.Fatal("expected proxy dial path") } } func TestDialTLSViaProxy_RejectsUnsupportedScheme(t *testing.T) { _, err := dialTLSViaProxy("imap.example.com:993", "ftp://proxy:21", nil) if err == nil || !strings.Contains(err.Error(), "unsupported proxy scheme") { t.Fatalf("expected unsupported scheme error, got %v", err) } } func TestDialTLSViaProxy_RejectsMalformedURL(t *testing.T) { _, err := dialTLSViaProxy("imap.example.com:993", "://bad", nil) if err == nil { t.Fatal("expected error for malformed URL") } } // --- SOCKS5 path --- type fakeContextDialer struct{} func (fakeContextDialer) Dial(network, address string) (net.Conn, error) { return nil, errors.New("used dial") } func (fakeContextDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { if _, ok := ctx.Deadline(); !ok { return nil, errors.New("missing deadline") } return nil, errors.New("used dialcontext") } func TestDialTLSViaProxy_SOCKS5_UsesDialContextWithTimeout(t *testing.T) { orig := socks5DialerFactory t.Cleanup(func() { socks5DialerFactory = orig }) socks5DialerFactory = func(network, address string, auth *proxy.Auth, forward proxy.Dialer) (proxy.Dialer, error) { if _, ok := forward.(proxy.ContextDialer); !ok { t.Fatalf("expected ContextDialer, got %T", forward) } return fakeContextDialer{}, nil } _, err := dialTLSViaProxy("imap.example.com:993", "socks5://127.0.0.1:1080", nil) if err == nil || !strings.Contains(err.Error(), "proxy dial: used dialcontext") { t.Fatalf("expected DialContext path, got %v", err) } } type trackingConn struct { deadline time.Time } func (c *trackingConn) Read(b []byte) (int, error) { return 0, io.EOF } func (c *trackingConn) Write(b []byte) (int, error) { return 0, io.EOF } func (c *trackingConn) Close() error { return nil } func (c *trackingConn) LocalAddr() net.Addr { return dummyAddr("local") } func (c *trackingConn) RemoteAddr() net.Addr { return dummyAddr("remote") } func (c *trackingConn) SetDeadline(t time.Time) error { c.deadline = t return nil } func (c *trackingConn) SetReadDeadline(t time.Time) error { return nil } func (c *trackingConn) SetWriteDeadline(t time.Time) error { return nil } type dummyAddr string func (a dummyAddr) Network() string { return "tcp" } func (a dummyAddr) String() string { return string(a) } type contextConnDialer struct { conn *trackingConn } func (d contextConnDialer) Dial(network, address string) (net.Conn, error) { return nil, errors.New("used dial") } func (d contextConnDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { return d.conn, nil } func TestDialTLSViaProxy_SOCKS5_SetsDeadlineForTLSHandshake(t *testing.T) { orig := socks5DialerFactory t.Cleanup(func() { socks5DialerFactory = orig }) conn := &trackingConn{} socks5DialerFactory = func(network, address string, auth *proxy.Auth, forward proxy.Dialer) (proxy.Dialer, error) { return contextConnDialer{conn: conn}, nil } _, err := dialTLSViaProxy("imap.example.com:993", "socks5://127.0.0.1:1080", nil) if err == nil || !strings.Contains(err.Error(), "tls handshake") { t.Fatalf("expected tls handshake error, got %v", err) } if conn.deadline.IsZero() { t.Fatal("expected TLS handshake deadline to be set") } } // --- HTTP CONNECT path --- func startFakeProxy(t *testing.T, handler func(host string) bool) (addr string, cleanup func()) { t.Helper() ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } done := make(chan struct{}) go func() { defer close(done) for { conn, err := ln.Accept() if err != nil { return } go func(c net.Conn) { defer c.Close() br := bufio.NewReader(c) req, err := http.ReadRequest(br) if err != nil { return } if handler(req.URL.Host) { fmt.Fprintf(c, "HTTP/1.1 200 OK\r\n\r\n") } else { fmt.Fprintf(c, "HTTP/1.1 403 Forbidden\r\n\r\n") } }(conn) } }() return ln.Addr().String(), func() { ln.Close() <-done } } func TestDialTLSViaProxy_HTTPConnect_ReachesTLSHandshake(t *testing.T) { addr, cleanup := startFakeProxy(t, func(host string) bool { return true }) defer cleanup() _, err := dialTLSViaProxy("imap.example.com:993", fmt.Sprintf("http://%s", addr), nil) if err == nil || !strings.Contains(err.Error(), "tls handshake") { t.Fatalf("expected tls handshake error (CONNECT succeeded), got %v", err) } } func TestDialTLSViaProxy_HTTPConnect_RejectsOnProxyFailure(t *testing.T) { addr, cleanup := startFakeProxy(t, func(host string) bool { return false }) defer cleanup() _, err := dialTLSViaProxy("imap.example.com:993", fmt.Sprintf("http://%s", addr), nil) if err == nil || !strings.Contains(err.Error(), "proxy connect failed") { t.Fatalf("expected proxy connect failed error, got %v", err) } } func TestDialTLSViaProxy_HTTPConnect_SendsProxyAuth(t *testing.T) { var gotAuth string ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } go func() { defer ln.Close() conn, err := ln.Accept() if err != nil { return } defer conn.Close() br := bufio.NewReader(conn) req, err := http.ReadRequest(br) if err != nil { return } gotAuth = req.Header.Get("Proxy-Authorization") fmt.Fprintf(conn, "HTTP/1.1 200 OK\r\n\r\n") }() _, _ = dialTLSViaProxy("imap.example.com:993", fmt.Sprintf("http://testuser:testpass@%s", ln.Addr().String()), nil) expected := "Basic " + base64.StdEncoding.EncodeToString([]byte("testuser:testpass")) if gotAuth != expected { t.Fatalf("expected auth %q, got %q", expected, gotAuth) } } func TestDialTLSViaProxy_HTTPConnect_DialError(t *testing.T) { origDial := dialToProxy t.Cleanup(func() { dialToProxy = origDial }) dialToProxy = func(ctx context.Context, network, addr string) (net.Conn, error) { return nil, errors.New("dial refused") } _, err := dialTLSViaProxy("imap.example.com:993", "http://127.0.0.1:1", nil) if err == nil || !strings.Contains(err.Error(), "dial proxy") { t.Fatalf("expected dial proxy error, got %v", err) } } // --- HTTPS CONNECT path --- func generateTestCert(t *testing.T) tls.Certificate { t.Helper() key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) if err != nil { t.Fatalf("generate key: %v", err) } tmpl := &x509.Certificate{ SerialNumber: big.NewInt(1), NotBefore: time.Now(), NotAfter: time.Now().Add(time.Hour), KeyUsage: x509.KeyUsageDigitalSignature, DNSNames: []string{"127.0.0.1"}, IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, } certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key) if err != nil { t.Fatalf("create cert: %v", err) } cert, err := x509.ParseCertificate(certDER) if err != nil { t.Fatalf("parse cert: %v", err) } return tls.Certificate{ Certificate: [][]byte{certDER}, PrivateKey: key, Leaf: cert, } } func TestDialTLSViaProxy_HTTPSConnect_ReachesTLSHandshake(t *testing.T) { cert := generateTestCert(t) ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } tlsLn := tls.NewListener(ln, &tls.Config{ Certificates: []tls.Certificate{cert}, }) done := make(chan struct{}) go func() { defer close(done) for { conn, err := tlsLn.Accept() if err != nil { return } go func(c net.Conn) { defer c.Close() br := bufio.NewReader(c) req, err := http.ReadRequest(br) if err != nil { return } if req.URL.Host != "" { fmt.Fprintf(c, "HTTP/1.1 200 OK\r\n\r\n") } }(conn) } }() defer func() { tlsLn.Close() <-done }() // Override dialToProxy to return raw TCP, but inject InsecureSkipVerify // so the TLS handshake to proxy succeeds with the test cert. origDial := dialToProxy t.Cleanup(func() { dialToProxy = origDial }) origProxyTLS := proxyTLSConfigForTest t.Cleanup(func() { proxyTLSConfigForTest = origProxyTLS }) proxyTLSConfigForTest = func(host string) *tls.Config { return &tls.Config{ServerName: host, InsecureSkipVerify: true} } _, err = dialTLSViaProxy("imap.example.com:993", fmt.Sprintf("https://%s", ln.Addr().String()), nil) if err == nil || !strings.Contains(err.Error(), "tls handshake") { t.Fatalf("expected tls handshake error (HTTPS CONNECT succeeded), got %v", err) } }