diff --git a/component/sniffer/dispatcher.go b/component/sniffer/dispatcher.go index 34a4512d..4f1aedda 100644 --- a/component/sniffer/dispatcher.go +++ b/component/sniffer/dispatcher.go @@ -2,7 +2,6 @@ package sniffer import ( "errors" - "net" "net/netip" "time" @@ -227,11 +226,9 @@ func (sd *Dispatcher) sniffDomain(conn *N.BufferedConn, metadata *C.Metadata) (s _, err := conn.Peek(1) _ = conn.SetReadDeadline(time.Time{}) if err != nil { - if _, ok := err.(*net.OpError); ok { - sd.cacheSniffFailed(metadata) - log.Errorln("[Sniffer] [%s] may not have any sent data, Consider adding skip", metadata.DstIP) - _ = conn.Close() - } + // The caller owns failure accounting and the connection lifetime. No + // initial data can be valid for a server-first protocol, so sniffing must + // not close the connection merely because this deadline expired. log.Debugln("[Sniffer] [%s] the data length not enough, error: %v", metadata.DstIP, err) return "", SnifferConfig{}, err } diff --git a/component/sniffer/dispatcher_test.go b/component/sniffer/dispatcher_test.go index 5ff342ab..87b9b98b 100644 --- a/component/sniffer/dispatcher_test.go +++ b/component/sniffer/dispatcher_test.go @@ -90,6 +90,10 @@ type stubSniffer struct { reply func(data []byte) (string, error) } +type domainMatcherFunc func(domain string) bool + +func (f domainMatcherFunc) MatchDomain(domain string) bool { return f(domain) } + var _ sniffer.Sniffer = (*stubSniffer)(nil) func (s *stubSniffer) SupportNetwork() C.NetWork { return s.network } @@ -345,6 +349,65 @@ func TestDispatcherMultipleSniffers(t *testing.T) { assert.False(t, cached) }) + t.Run("keeps the connection open and caches an initial read failure once", func(t *testing.T) { + waiting := &stubSniffer{reply: func(data []byte) (string, error) { + return "", needAtLeast(len(data) + 1) + }} + sd, err := NewDispatcher(&Config{Enable: true, ParsePureIp: true}) + assert.NoError(t, err) + sd.sniffers = []configuredSniffer{{Sniffer: waiting}} + raw := &chunkedConn{ + t: t, + readErr: &net.OpError{Op: "read", Net: "tcp", Err: io.ErrNoProgress}, + } + t.Cleanup(raw.stopTimer) + metadata := &C.Metadata{ + NetWork: C.TCP, + DstIP: netip.MustParseAddr("192.0.2.1"), + DstPort: 80, + } + + sniffed := sd.TCPSniff(N.NewBufferedConn(raw), metadata) + + assert.False(t, sniffed) + assert.False(t, raw.closed) + assert.Empty(t, waiting.seen) + failures, cached := sd.skipList.Get(metadata.AddrPort()) + assert.True(t, cached) + assert.Equal(t, uint8(1), failures) + assert.True(t, raw.deadline.IsZero(), "read deadline was not cleared") + }) + + t.Run("does not cache a forced initial read failure", func(t *testing.T) { + waiting := &stubSniffer{reply: func(data []byte) (string, error) { + return "", needAtLeast(len(data) + 1) + }} + sd, err := NewDispatcher(&Config{ + Enable: true, + ForceDomain: []C.DomainMatcher{domainMatcherFunc(func(string) bool { return true })}, + }) + assert.NoError(t, err) + sd.sniffers = []configuredSniffer{{Sniffer: waiting}} + raw := &chunkedConn{ + t: t, + readErr: &net.OpError{Op: "read", Net: "tcp", Err: io.ErrNoProgress}, + } + t.Cleanup(raw.stopTimer) + metadata := &C.Metadata{ + NetWork: C.TCP, + Host: "forced.example", + DstIP: netip.MustParseAddr("192.0.2.1"), + DstPort: 80, + } + + sniffed := sd.TCPSniff(N.NewBufferedConn(raw), metadata) + + assert.False(t, sniffed) + assert.False(t, raw.closed) + _, cached := sd.skipList.Get(metadata.AddrPort()) + assert.False(t, cached) + }) + t.Run("accepts all-network sniffer", func(t *testing.T) { allNetwork := &stubSniffer{ network: C.ALLNet,