From d2fd2ab794a11ca6d22e45395bf0ee7799cf3d01 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Fri, 17 Apr 2026 16:51:53 +0800 Subject: [PATCH] Fix tls-spoof --- common/tls/client.go | 10 +- common/tls/client_test.go | 154 +++++++++++ common/tls/std_client.go | 2 +- common/tls/utls_client.go | 2 +- common/tls/utls_client_test.go | 73 ++++++ common/tlsspoof/client_hello.go | 103 ++------ common/tlsspoof/client_hello_test.go | 110 ++++---- common/tlsspoof/conn_test.go | 267 +++++++++++++++++++- common/tlsspoof/integration_test.go | 33 ++- common/tlsspoof/integration_tls_test.go | 118 +++++++++ common/tlsspoof/integration_unix_test.go | 97 +++++-- common/tlsspoof/integration_windows_test.go | 27 +- common/tlsspoof/packet_test.go | 59 +++++ common/tlsspoof/raw_darwin.go | 41 ++- common/tlsspoof/raw_linux.go | 24 +- common/tlsspoof/raw_stub.go | 2 +- common/tlsspoof/raw_windows.go | 32 ++- common/tlsspoof/spoof.go | 56 ++-- common/windivert/handle_windows.go | 6 +- common/windivert/handle_windows_test.go | 3 + common/windivert/windivert.go | 9 +- 21 files changed, 985 insertions(+), 243 deletions(-) create mode 100644 common/tls/client_test.go create mode 100644 common/tls/utls_client_test.go create mode 100644 common/tlsspoof/integration_tls_test.go diff --git a/common/tls/client.go b/common/tls/client.go index 35c628c11..b5b975bf2 100644 --- a/common/tls/client.go +++ b/common/tls/client.go @@ -6,6 +6,7 @@ import ( "errors" "net" "os" + "strings" "github.com/sagernet/sing-box/common/badtls" "github.com/sagernet/sing-box/common/tlsspoof" @@ -33,6 +34,9 @@ func parseTLSSpoofOptions(serverName string, options option.OutboundTLSOptions) if options.DisableSNI || serverName == "" || M.ParseAddr(serverName).IsValid() { return "", 0, E.New("`spoof` requires TLS ClientHello with SNI") } + if strings.EqualFold(options.Spoof, serverName) { + return "", 0, E.New("`spoof` must differ from `server_name`") + } method, err := tlsspoof.ParseMethod(options.SpoofMethod) if err != nil { return "", 0, err @@ -44,11 +48,7 @@ func applyTLSSpoof(conn net.Conn, spoof string, method tlsspoof.Method) (net.Con if spoof == "" { return conn, nil } - spoofer, err := tlsspoof.NewSpoofer(conn, method) - if err != nil { - return nil, err - } - return tlsspoof.NewConn(conn, spoofer, spoof), nil + return tlsspoof.NewConn(conn, method, spoof) } func NewDialerFromOptions(ctx context.Context, logger logger.ContextLogger, dialer N.Dialer, serverAddress string, options option.OutboundTLSOptions) (N.Dialer, error) { diff --git a/common/tls/client_test.go b/common/tls/client_test.go new file mode 100644 index 000000000..5bc939e29 --- /dev/null +++ b/common/tls/client_test.go @@ -0,0 +1,154 @@ +package tls + +import ( + "context" + "crypto/tls" + "net" + "testing" + + tf "github.com/sagernet/sing-box/common/tlsfragment" + "github.com/sagernet/sing-box/common/tlsspoof" + "github.com/sagernet/sing-box/option" + + "github.com/stretchr/testify/require" +) + +func TestParseTLSSpoofOptions_Disabled(t *testing.T) { + t.Parallel() + spoof, method, err := parseTLSSpoofOptions("example.com", option.OutboundTLSOptions{}) + require.NoError(t, err) + require.Empty(t, spoof) + require.Equal(t, tlsspoof.MethodWrongSequence, method) +} + +func TestParseTLSSpoofOptions_MethodWithoutSpoof(t *testing.T) { + t.Parallel() + _, _, err := parseTLSSpoofOptions("example.com", option.OutboundTLSOptions{ + SpoofMethod: tlsspoof.MethodNameWrongChecksum, + }) + require.Error(t, err) +} + +func TestParseTLSSpoofOptions_IPLiteralRejected(t *testing.T) { + t.Parallel() + _, _, err := parseTLSSpoofOptions("1.2.3.4", option.OutboundTLSOptions{ + Spoof: "example.com", + }) + require.Error(t, err) +} + +func TestParseTLSSpoofOptions_EmptyServerNameRejected(t *testing.T) { + t.Parallel() + _, _, err := parseTLSSpoofOptions("", option.OutboundTLSOptions{ + Spoof: "example.com", + }) + require.Error(t, err) +} + +func TestParseTLSSpoofOptions_DisableSNIRejected(t *testing.T) { + t.Parallel() + _, _, err := parseTLSSpoofOptions("example.com", option.OutboundTLSOptions{ + Spoof: "decoy.com", + DisableSNI: true, + }) + require.Error(t, err) +} + +// TestParseTLSSpoofOptions_RejectsSameSNI is the primary regression test for +// the "spoofed packet contains the original SNI" bug report: when a user +// configures spoof equal to server_name, the rewriter produces a byte-identical +// record, so the fake and real ClientHellos on the wire look the same. Reject +// at parse time. +func TestParseTLSSpoofOptions_RejectsSameSNI(t *testing.T) { + t.Parallel() + _, _, err := parseTLSSpoofOptions("example.com", option.OutboundTLSOptions{ + Spoof: "example.com", + }) + require.Error(t, err) + + _, _, err = parseTLSSpoofOptions("example.com", option.OutboundTLSOptions{ + Spoof: "EXAMPLE.com", + }) + require.Error(t, err, "comparison must be case-insensitive") +} + +func TestParseTLSSpoofOptions_UnknownMethodRejected(t *testing.T) { + t.Parallel() + _, _, err := parseTLSSpoofOptions("example.com", option.OutboundTLSOptions{ + Spoof: "decoy.com", + SpoofMethod: "nonsense", + }) + require.Error(t, err) +} + +func TestParseTLSSpoofOptions_DistinctSNIAccepted(t *testing.T) { + t.Parallel() + if !tlsspoof.PlatformSupported { + t.Skip("tlsspoof not supported on this platform") + } + spoof, method, err := parseTLSSpoofOptions("example.com", option.OutboundTLSOptions{ + Spoof: "decoy.com", + SpoofMethod: tlsspoof.MethodNameWrongSequence, + }) + require.NoError(t, err) + require.Equal(t, "decoy.com", spoof) + require.Equal(t, tlsspoof.MethodWrongSequence, method) +} + +// The following tests guard the wrap gate in STDClientConfig.Client(): +// tf.Conn must wrap the underlying connection whenever either `fragment` or +// `record_fragment` is set, so that TLS fragmentation coexists with features +// like tls_spoof that layer on top of tf.Conn. + +func newSTDClientConfigForGateTest(fragment, recordFragment bool) *STDClientConfig { + return &STDClientConfig{ + ctx: context.Background(), + config: &tls.Config{ServerName: "example.com", InsecureSkipVerify: true}, + fragment: fragment, + recordFragment: recordFragment, + } +} + +func TestSTDClient_Client_NoFragment_DoesNotWrap(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newSTDClientConfigForGateTest(false, false).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.False(t, isTF, "no fragment flags: must not wrap with tf.Conn") +} + +func TestSTDClient_Client_FragmentOnly_Wraps(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newSTDClientConfigForGateTest(true, false).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.True(t, isTF, "fragment=true: must wrap with tf.Conn") +} + +func TestSTDClient_Client_RecordFragmentOnly_Wraps(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newSTDClientConfigForGateTest(false, true).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.True(t, isTF, "record_fragment=true: must wrap with tf.Conn") +} + +func TestSTDClient_Client_BothFragment_Wraps(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newSTDClientConfigForGateTest(true, true).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.True(t, isTF, "both fragment flags: must wrap with tf.Conn") +} diff --git a/common/tls/std_client.go b/common/tls/std_client.go index f38981c68..031a256f7 100644 --- a/common/tls/std_client.go +++ b/common/tls/std_client.go @@ -75,7 +75,7 @@ func (c *STDClientConfig) STDConfig() (*STDConfig, error) { } func (c *STDClientConfig) Client(conn net.Conn) (Conn, error) { - if c.recordFragment { + if c.fragment || c.recordFragment { conn = tf.NewConn(conn, c.ctx, c.fragment, c.recordFragment, c.fragmentFallbackDelay) } conn, err := applyTLSSpoof(conn, c.spoof, c.spoofMethod) diff --git a/common/tls/utls_client.go b/common/tls/utls_client.go index a8b91973c..1cc41554f 100644 --- a/common/tls/utls_client.go +++ b/common/tls/utls_client.go @@ -83,7 +83,7 @@ func (c *UTLSClientConfig) STDConfig() (*STDConfig, error) { } func (c *UTLSClientConfig) Client(conn net.Conn) (Conn, error) { - if c.recordFragment { + if c.fragment || c.recordFragment { conn = tf.NewConn(conn, c.ctx, c.fragment, c.recordFragment, c.fragmentFallbackDelay) } conn, err := applyTLSSpoof(conn, c.spoof, c.spoofMethod) diff --git a/common/tls/utls_client_test.go b/common/tls/utls_client_test.go new file mode 100644 index 000000000..48c1e327e --- /dev/null +++ b/common/tls/utls_client_test.go @@ -0,0 +1,73 @@ +//go:build with_utls + +package tls + +import ( + "context" + "net" + "testing" + + tf "github.com/sagernet/sing-box/common/tlsfragment" + + utls "github.com/metacubex/utls" + "github.com/stretchr/testify/require" +) + +// Guards the wrap gate in UTLSClientConfig.Client(): tf.Conn must wrap the +// underlying connection whenever either `fragment` or `record_fragment` is +// set. Mirrors the STDClientConfig gate tests to keep both code paths in +// lockstep. + +func newUTLSClientConfigForGateTest(fragment, recordFragment bool) *UTLSClientConfig { + return &UTLSClientConfig{ + ctx: context.Background(), + config: &utls.Config{ServerName: "example.com", InsecureSkipVerify: true}, + id: utls.HelloChrome_Auto, + fragment: fragment, + recordFragment: recordFragment, + } +} + +func TestUTLSClient_Client_NoFragment_DoesNotWrap(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newUTLSClientConfigForGateTest(false, false).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.False(t, isTF, "no fragment flags: must not wrap with tf.Conn") +} + +func TestUTLSClient_Client_FragmentOnly_Wraps(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newUTLSClientConfigForGateTest(true, false).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.True(t, isTF, "fragment=true: must wrap with tf.Conn") +} + +func TestUTLSClient_Client_RecordFragmentOnly_Wraps(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newUTLSClientConfigForGateTest(false, true).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.True(t, isTF, "record_fragment=true: must wrap with tf.Conn") +} + +func TestUTLSClient_Client_BothFragment_Wraps(t *testing.T) { + t.Parallel() + client, server := net.Pipe() + defer client.Close() + defer server.Close() + wrapped, err := newUTLSClientConfigForGateTest(true, true).Client(client) + require.NoError(t, err) + _, isTF := wrapped.NetConn().(*tf.Conn) + require.True(t, isTF, "both fragment flags: must wrap with tf.Conn") +} diff --git a/common/tlsspoof/client_hello.go b/common/tlsspoof/client_hello.go index 0ca7c5a9f..abdfa3175 100644 --- a/common/tlsspoof/client_hello.go +++ b/common/tlsspoof/client_hello.go @@ -1,86 +1,37 @@ package tlsspoof import ( - "encoding/binary" + "bytes" + "context" + "crypto/tls" - tf "github.com/sagernet/sing-box/common/tlsfragment" + "github.com/sagernet/sing/common/bufio" E "github.com/sagernet/sing/common/exceptions" ) -const ( - recordLengthOffset = 3 - handshakeLengthOffset = 6 -) - -// server_name extension layout (RFC 6066 §3). Offsets are relative to the -// SNI host name (index returned by the parser): -// -// ... uint16 extension_type = 0x0000 (host_name - 9) -// ... uint16 extension_data_length (host_name - 7) -// ... uint16 server_name_list_length (host_name - 5) -// ... uint8 name_type = host_name (host_name - 3) -// ... uint16 host_name_length (host_name - 2) -// sni host_name (host_name) -const ( - extensionDataLengthOffsetFromSNI = -7 - listLengthOffsetFromSNI = -5 - hostNameLengthOffsetFromSNI = -2 -) - -func rewriteSNI(record []byte, fakeSNI string) ([]byte, error) { - if len(fakeSNI) > 0xFFFF { - return nil, E.New("fake SNI too long: ", len(fakeSNI), " bytes") +// buildFakeClientHello drives crypto/tls against a write-only in-memory conn +// to capture a generated ClientHello. CurvePreferences pins classical groups +// to suppress Go's default X25519MLKEM768 hybrid key share; without this the +// post-quantum public key alone (~1184 bytes) pushes the record past one MSS, +// and middleboxes do not reassemble fragmented ClientHellos. The handshake +// error is discarded because the stub conn's Read returns immediately. +func buildFakeClientHello(sni string) ([]byte, error) { + if sni == "" { + return nil, E.New("empty sni") } - serverName := tf.IndexTLSServerName(record) - if serverName == nil { - return nil, E.New("not a ClientHello with SNI") + var buf bytes.Buffer + tlsConn := tls.Client(bufio.NewWriteOnlyConn(&buf), &tls.Config{ + ServerName: sni, + // Order matches what browsers advertised before post-quantum. + CurvePreferences: []tls.CurveID{tls.X25519, tls.CurveP256, tls.CurveP384}, + MinVersion: tls.VersionTLS12, + MaxVersion: tls.VersionTLS13, + NextProtos: []string{"h2", "http/1.1"}, + InsecureSkipVerify: true, + }) + _ = tlsConn.HandshakeContext(context.Background()) + if buf.Len() == 0 { + return nil, E.New("tls ClientHello not produced") } - - delta := len(fakeSNI) - serverName.Length - out := make([]byte, len(record)+delta) - copy(out, record[:serverName.Index]) - copy(out[serverName.Index:], fakeSNI) - copy(out[serverName.Index+len(fakeSNI):], record[serverName.Index+serverName.Length:]) - - err := patchUint16(out, recordLengthOffset, delta) - if err != nil { - return nil, E.Cause(err, "patch record length") - } - err = patchUint24(out, handshakeLengthOffset, delta) - if err != nil { - return nil, E.Cause(err, "patch handshake length") - } - for _, off := range []int{ - serverName.ExtensionsListLengthIndex, - serverName.Index + extensionDataLengthOffsetFromSNI, - serverName.Index + listLengthOffsetFromSNI, - serverName.Index + hostNameLengthOffsetFromSNI, - } { - err = patchUint16(out, off, delta) - if err != nil { - return nil, E.Cause(err, "patch length at offset ", off) - } - } - return out, nil -} - -func patchUint16(data []byte, offset, delta int) error { - patched := int(binary.BigEndian.Uint16(data[offset:])) + delta - if patched < 0 || patched > 0xFFFF { - return E.New("uint16 out of range: ", patched) - } - binary.BigEndian.PutUint16(data[offset:], uint16(patched)) - return nil -} - -func patchUint24(data []byte, offset, delta int) error { - original := int(data[offset])<<16 | int(data[offset+1])<<8 | int(data[offset+2]) - patched := original + delta - if patched < 0 || patched > 0xFFFFFF { - return E.New("uint24 out of range: ", patched) - } - data[offset] = byte(patched >> 16) - data[offset+1] = byte(patched >> 8) - data[offset+2] = byte(patched) - return nil + return buf.Bytes(), nil } diff --git a/common/tlsspoof/client_hello_test.go b/common/tlsspoof/client_hello_test.go index 746d0482a..3eb7a2e04 100644 --- a/common/tlsspoof/client_hello_test.go +++ b/common/tlsspoof/client_hello_test.go @@ -1,8 +1,9 @@ package tlsspoof import ( + "bytes" "encoding/binary" - "encoding/hex" + "strings" "testing" tf "github.com/sagernet/sing-box/common/tlsfragment" @@ -10,70 +11,73 @@ import ( "github.com/stretchr/testify/require" ) -// realClientHello is a captured Chrome ClientHello for github.com, -// reused from common/tlsfragment/index_test.go. -const realClientHello = "16030105f8010005f403036e35de7389a679c54029cf452611f2211c70d9ac3897271de589ab6155f8e4ab20637d225f1ef969ad87ed78bfb9d171300bcb1703b6f314ccefb964f79b7d0961002a0a0a130213031301c02cc02bcca9c030c02fcca8c00ac009c014c013009d009c0035002fc008c012000a01000581baba00000000000f000d00000a6769746875622e636f6d00170000ff01000100000a000e000c3a3a11ec001d001700180019000b000201000010000e000c02683208687474702f312e31000500050100000000000d00160014040308040401050308050805050108060601020100120000003304ef04ed3a3a00010011ec04c0aeb2250c092a3463161cccb29d9183331a424964248579507ed23a180b0ceab2a5f5d9ce41547e497a89055471ea572867ba3a1fc3c9e45025274a20f60c6b60e62476b6afed0403af59ab83660ef4112ae20386a602010d0a5d454c0ed34c84ed4423e750213e6a2baab1bf9c4367a6007ab40a33d95220c2dcaa44f257024a5626b545db0510f4311b1a60714154909c6a61fdfca011fb2626d657aeb6070bf078508babe3b584555013e34acc56198ed4663742b3155a664a9901794c4586820a7dc162c01827291f3792e1237f801a8d1ef096013c181c4a58d2f6859ba75022d18cc4418bd4f351d5c18f83a58857d05af860c4b9ac018a5b63f17184e591532c6bc2cf2215d4a282c8a8a4f6f7aee110422c8bc9ebd3b1d609c568523aaae555db320e6c269473d87af38c256cbb9febc20aea6380c32a8916f7a373c8b1e37554e3260bf6621f6b804ee80b3c516b1d01985bf4c603b6daa9a5991de6a7a29f3a7122b8afb843a7660110fce62b43c615f5bcc2db688ba012649c0952b0a2c031e732d2b454c6b2968683cb8d244be2c9a7fa163222979eaf92722b92b862d81a3d94450c2b60c318421ebb4307c42d1f0473592a5c30e42039cc68cda9721e61aa63f49def17c15221680ed444896340133bbee67556f56b9f9d78a4df715f926a12add0cc9c862e46ea8b7316ae468282c18601b2771c9c9322f982228cf93effaacd3f80cbd12bce5fc36f56e2a3caf91e578a5fae00c9b23a8ed1a66764f4433c3628a70b8f0a6196adc60a4cb4226f07ba4c6b363fe9065563bfc1347452946386bab488686e837ab979c64f9047417fca635fe1bb4f074f256cc8af837c7b455e280426547755af90a61640169ef180aea3a77e662bb6dac1b6c3696027129b1a5edf495314e9c7f4b6110e16378ec893fa24642330a40aba1a85326101acb97c620fd8d71389e69eaed7bdb01bbe1fd428d66191150c7b2cd1ad4257391676a82ba8ce07fb2667c3b289f159003a7c7bc31d361b7b7f49a802961739d950dfcc0fa1c7abce5abdd2245101da391151490862028110465950b9e9c03d08a90998ab83267838d2e74a0593bc81f74cdf734519a05b351c0e5488c68dd810e6e9142ccc1e2f4a7f464297eb340e27acc6b9d64e12e38cce8492b3d939140b5a9e149a75597f10a23874c84323a07cdd657274378f887c85c4259b9c04cd33ba58ed630ef2a744f8e19dd34843dff331d2a6be7e2332c599289cd248a611c73d7481cd4a9bd43449a3836f14b2af18a1739e17999e4c67e85cc5bcecabb14185e5bcaff3c96098f03dc5aba819f29587758f49f940585354a2a780830528d68ccd166920dadcaa25cab5fc1907272a826aba3f08bc6b88757776812ecb6c7cec69a223ec0a13a7b62a2349a0f63ed7a27a3b15ba21d71fe6864ec6e089ae17cadd433fa3138f7ee24353c11365818f8fc34f43a05542d18efaac24bfccc1f748a0cc1a67ad379468b76fd34973dba785f5c91d618333cd810fe0700d1bbc8422029782628070a624c52c5309a4a64d625b11f8033ab28df34a1add297517fcc06b92b6817b3c5144438cf260867c57bde68c8c4b82e6a135ef676a52fbae5708002a404e6189a60e2836de565ad1b29e3819e5ed49f6810bcb28e1bd6de57306f94b79d9dae1cc4624d2a068499beef81cd5fe4b76dcbfff2a2008001d002001976128c6d5a934533f28b9914d2480aab2a8c1ab03d212529ce8b27640a716002d00020101002b000706caca03040303001b00030200015a5a000100" +// x25519MLKEM768 is the IANA code point for the post-quantum hybrid named +// group (0x11EC). The fake ClientHello must never carry it — its 1184-byte +// key share is the reason kernel-generated ClientHellos exceed one MSS, and +// the reason this builder has to force CurvePreferences. +const x25519MLKEM768 uint16 = 0x11EC -func decodeClientHello(t *testing.T) []byte { - t.Helper() - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) - return payload -} - -func assertConsistent(t *testing.T, payload []byte, expectedSNI string) { - t.Helper() - serverName := tf.IndexTLSServerName(payload) - require.NotNil(t, serverName, "parser should find SNI in rewritten payload") - require.Equal(t, expectedSNI, serverName.ServerName) - require.Equal(t, expectedSNI, string(payload[serverName.Index:serverName.Index+serverName.Length])) - // Record length must equal len(payload) - 5. - recordLen := binary.BigEndian.Uint16(payload[3:5]) - require.Equal(t, len(payload)-5, int(recordLen), "record length must equal payload - 5") - // Handshake length must equal len(payload) - 5 - 4. - handshakeLen := int(payload[6])<<16 | int(payload[7])<<8 | int(payload[8]) - require.Equal(t, len(payload)-5-4, handshakeLen, "handshake length must equal payload - 9") -} - -func TestRewriteSNI_ShorterReplacement(t *testing.T) { +func TestBuildFakeClientHello_ParsesWithSNI(t *testing.T) { t.Parallel() - payload := decodeClientHello(t) - out, err := rewriteSNI(payload, "a.io") + record, err := buildFakeClientHello("example.com") require.NoError(t, err) - require.Len(t, out, len(payload)-6) // original "github.com" is 10 bytes, "a.io" is 4 bytes. - assertConsistent(t, out, "a.io") + + serverName := tf.IndexTLSServerName(record) + require.NotNil(t, serverName, "output must parse as a ClientHello") + require.Equal(t, "example.com", serverName.ServerName) + + recordLen := binary.BigEndian.Uint16(record[3:5]) + require.Equal(t, len(record)-5, int(recordLen), + "record length header must match on-wire record size") + handshakeLen := int(record[6])<<16 | int(record[7])<<8 | int(record[8]) + require.Equal(t, len(record)-5-4, handshakeLen, + "handshake length header must match handshake body size") } -func TestRewriteSNI_SameLengthReplacement(t *testing.T) { +// TestBuildFakeClientHello_FitsOneSegment is the regression guard for the +// whole point of the rewrite: the fake must never need fragmenting on a +// standard 1500-byte path MTU. 1200 leaves ~260 bytes for IP+TCP headers and +// a generous safety margin — the X25519MLKEM768 ClientHello this replaces +// hit ~1400+. +func TestBuildFakeClientHello_FitsOneSegment(t *testing.T) { t.Parallel() - payload := decodeClientHello(t) - out, err := rewriteSNI(payload, "example.co") - require.NoError(t, err) - require.Len(t, out, len(payload)) - assertConsistent(t, out, "example.co") + for _, sni := range []string{"a.io", "example.com", strings.Repeat("a", 253)} { + record, err := buildFakeClientHello(sni) + require.NoError(t, err, "sni=%q", sni) + require.Less(t, len(record), 1200, "sni=%q built %d bytes", sni, len(record)) + } } -func TestRewriteSNI_LongerReplacement(t *testing.T) { +// TestBuildFakeClientHello_NoPostQuantumKeyShare catches regressions that +// would accidentally pull an X25519MLKEM768 key share (the reason the prior +// implementation had to fragment) back into the fake — e.g. if CurvePreferences +// stopped being respected by a future Go version. +func TestBuildFakeClientHello_NoPostQuantumKeyShare(t *testing.T) { t.Parallel() - payload := decodeClientHello(t) - out, err := rewriteSNI(payload, "letsencrypt.org") + record, err := buildFakeClientHello("example.com") require.NoError(t, err) - require.Len(t, out, len(payload)+5) // "letsencrypt.org" is 15, original 10, delta 5. - assertConsistent(t, out, "letsencrypt.org") + + var needle [2]byte + binary.BigEndian.PutUint16(needle[:], x25519MLKEM768) + require.False(t, bytes.Contains(record, needle[:]), + "output must not contain the X25519MLKEM768 code point (0x%04x)", x25519MLKEM768) } -func TestRewriteSNI_NoSNIReturnsError(t *testing.T) { +// TestBuildFakeClientHello_RandomizesPerCall ensures crypto/tls generates a +// fresh random + session_id + key_share on every call, as required to avoid +// trivial fingerprinting of the spoof. +func TestBuildFakeClientHello_RandomizesPerCall(t *testing.T) { t.Parallel() - // Truncated payload — not a valid ClientHello. - _, err := rewriteSNI([]byte{0x16, 0x03, 0x01, 0x00, 0x01, 0x01}, "x.com") + first, err := buildFakeClientHello("example.com") + require.NoError(t, err) + second, err := buildFakeClientHello("example.com") + require.NoError(t, err) + require.NotEqual(t, first, second, + "repeated calls must produce distinct bytes (random/session_id/key_share must vary)") +} + +func TestBuildFakeClientHello_RejectsEmpty(t *testing.T) { + t.Parallel() + _, err := buildFakeClientHello("") require.Error(t, err) } - -func TestRewriteSNI_DoesNotMutateInput(t *testing.T) { - t.Parallel() - payload := decodeClientHello(t) - original := append([]byte(nil), payload...) - _, err := rewriteSNI(payload, "letsencrypt.org") - require.NoError(t, err) - require.Equal(t, original, payload, "input payload must not be mutated") -} diff --git a/common/tlsspoof/conn_test.go b/common/tlsspoof/conn_test.go index 981f1a49c..b41cf5475 100644 --- a/common/tlsspoof/conn_test.go +++ b/common/tlsspoof/conn_test.go @@ -1,19 +1,36 @@ package tlsspoof import ( + "bytes" + "context" + "encoding/binary" "encoding/hex" "io" "net" "testing" + "time" tf "github.com/sagernet/sing-box/common/tlsfragment" "github.com/stretchr/testify/require" ) +// realClientHello is a captured Chrome ClientHello for github.com. Tests that +// stack tlsspoof.Conn on top of tf.Conn still need a parseable payload to +// exercise the fragment transform. +const realClientHello = "16030105f8010005f403036e35de7389a679c54029cf452611f2211c70d9ac3897271de589ab6155f8e4ab20637d225f1ef969ad87ed78bfb9d171300bcb1703b6f314ccefb964f79b7d0961002a0a0a130213031301c02cc02bcca9c030c02fcca8c00ac009c014c013009d009c0035002fc008c012000a01000581baba00000000000f000d00000a6769746875622e636f6d00170000ff01000100000a000e000c3a3a11ec001d001700180019000b000201000010000e000c02683208687474702f312e31000500050100000000000d00160014040308040401050308050805050108060601020100120000003304ef04ed3a3a00010011ec04c0aeb2250c092a3463161cccb29d9183331a424964248579507ed23a180b0ceab2a5f5d9ce41547e497a89055471ea572867ba3a1fc3c9e45025274a20f60c6b60e62476b6afed0403af59ab83660ef4112ae20386a602010d0a5d454c0ed34c84ed4423e750213e6a2baab1bf9c4367a6007ab40a33d95220c2dcaa44f257024a5626b545db0510f4311b1a60714154909c6a61fdfca011fb2626d657aeb6070bf078508babe3b584555013e34acc56198ed4663742b3155a664a9901794c4586820a7dc162c01827291f3792e1237f801a8d1ef096013c181c4a58d2f6859ba75022d18cc4418bd4f351d5c18f83a58857d05af860c4b9ac018a5b63f17184e591532c6bc2cf2215d4a282c8a8a4f6f7aee110422c8bc9ebd3b1d609c568523aaae555db320e6c269473d87af38c256cbb9febc20aea6380c32a8916f7a373c8b1e37554e3260bf6621f6b804ee80b3c516b1d01985bf4c603b6daa9a5991de6a7a29f3a7122b8afb843a7660110fce62b43c615f5bcc2db688ba012649c0952b0a2c031e732d2b454c6b2968683cb8d244be2c9a7fa163222979eaf92722b92b862d81a3d94450c2b60c318421ebb4307c42d1f0473592a5c30e42039cc68cda9721e61aa63f49def17c15221680ed444896340133bbee67556f56b9f9d78a4df715f926a12add0cc9c862e46ea8b7316ae468282c18601b2771c9c9322f982228cf93effaacd3f80cbd12bce5fc36f56e2a3caf91e578a5fae00c9b23a8ed1a66764f4433c3628a70b8f0a6196adc60a4cb4226f07ba4c6b363fe9065563bfc1347452946386bab488686e837ab979c64f9047417fca635fe1bb4f074f256cc8af837c7b455e280426547755af90a61640169ef180aea3a77e662bb6dac1b6c3696027129b1a5edf495314e9c7f4b6110e16378ec893fa24642330a40aba1a85326101acb97c620fd8d71389e69eaed7bdb01bbe1fd428d66191150c7b2cd1ad4257391676a82ba8ce07fb2667c3b289f159003a7c7bc31d361b7b7f49a802961739d950dfcc0fa1c7abce5abdd2245101da391151490862028110465950b9e9c03d08a90998ab83267838d2e74a0593bc81f74cdf734519a05b351c0e5488c68dd810e6e9142ccc1e2f4a7f464297eb340e27acc6b9d64e12e38cce8492b3d939140b5a9e149a75597f10a23874c84323a07cdd657274378f887c85c4259b9c04cd33ba58ed630ef2a744f8e19dd34843dff331d2a6be7e2332c599289cd248a611c73d7481cd4a9bd43449a3836f14b2af18a1739e17999e4c67e85cc5bcecabb14185e5bcaff3c96098f03dc5aba819f29587758f49f940585354a2a780830528d68ccd166920dadcaa25cab5fc1907272a826aba3f08bc6b88757776812ecb6c7cec69a223ec0a13a7b62a2349a0f63ed7a27a3b15ba21d71fe6864ec6e089ae17cadd433fa3138f7ee24353c11365818f8fc34f43a05542d18efaac24bfccc1f748a0cc1a67ad379468b76fd34973dba785f5c91d618333cd810fe0700d1bbc8422029782628070a624c52c5309a4a64d625b11f8033ab28df34a1add297517fcc06b92b6817b3c5144438cf260867c57bde68c8c4b82e6a135ef676a52fbae5708002a404e6189a60e2836de565ad1b29e3819e5ed49f6810bcb28e1bd6de57306f94b79d9dae1cc4624d2a068499beef81cd5fe4b76dcbfff2a2008001d002001976128c6d5a934533f28b9914d2480aab2a8c1ab03d212529ce8b27640a716002d00020101002b000706caca03040303001b00030200015a5a000100" + +func decodeClientHello(t *testing.T) []byte { + t.Helper() + payload, err := hex.DecodeString(realClientHello) + require.NoError(t, err) + return payload +} + type fakeSpoofer struct { injected [][]byte err error + closeErr error } func (f *fakeSpoofer) Inject(payload []byte) error { @@ -25,7 +42,7 @@ func (f *fakeSpoofer) Inject(payload []byte) error { } func (f *fakeSpoofer) Close() error { - return nil + return f.closeErr } func readAll(t *testing.T, conn net.Conn) []byte { @@ -37,12 +54,12 @@ func readAll(t *testing.T, conn net.Conn) []byte { func TestConn_Write_InjectsThenForwards(t *testing.T) { t.Parallel() - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) + payload := decodeClientHello(t) client, server := net.Pipe() spoofer := &fakeSpoofer{} - wrapped := NewConn(client, spoofer, "letsencrypt.org") + wrapped, err := newConn(client, spoofer, "letsencrypt.org") + require.NoError(t, err) serverRead := make(chan []byte, 1) go func() { @@ -66,12 +83,12 @@ func TestConn_Write_InjectsThenForwards(t *testing.T) { func TestConn_Write_SecondWriteDoesNotInject(t *testing.T) { t.Parallel() - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) + payload := decodeClientHello(t) client, server := net.Pipe() spoofer := &fakeSpoofer{} - wrapped := NewConn(client, spoofer, "letsencrypt.org") + wrapped, err := newConn(client, spoofer, "letsencrypt.org") + require.NoError(t, err) serverRead := make(chan []byte, 1) go func() { @@ -89,18 +106,244 @@ func TestConn_Write_SecondWriteDoesNotInject(t *testing.T) { require.Len(t, spoofer.injected, 1) } -func TestConn_Write_NonClientHelloReturnsError(t *testing.T) { +// TestConn_Write_SurfacesCloseError guards against the defer pattern silently +// dropping the spoofer's Close() error on the success path. +func TestConn_Write_SurfacesCloseError(t *testing.T) { + t.Parallel() + + client, server := net.Pipe() + defer client.Close() + defer server.Close() + spoofer := &fakeSpoofer{closeErr: errSpoofClose} + wrapped, err := newConn(client, spoofer, "letsencrypt.org") + require.NoError(t, err) + + go func() { _, _ = io.ReadAll(server) }() + + _, err = wrapped.Write([]byte("trigger inject")) + require.ErrorIs(t, err, errSpoofClose, + "Close() error must be wrapped into Write's return") +} + +func TestConn_NewConn_RejectsEmptySNI(t *testing.T) { t.Parallel() client, server := net.Pipe() defer client.Close() defer server.Close() + _, err := newConn(client, &fakeSpoofer{}, "") + require.Error(t, err, "empty SNI must fail at construction") +} +var errSpoofClose = errTest("spoof-close-failed") + +type errTest string + +func (e errTest) Error() string { return string(e) } + +// recordingConn intercepts each Write call so tests can assert how many +// downstream writes occurred and in what order with respect to spoof +// injection. It does not implement WithUpstream, so tf.Conn's +// N.UnwrapReader(conn).(*net.TCPConn) returns nil and fragment-mode falls +// back to its plain Write + time.Sleep path — which is what we want to +// exercise over a net.Pipe. +type recordingConn struct { + net.Conn + writes [][]byte + timeline *[]string +} + +func (c *recordingConn) Write(p []byte) (int, error) { + c.writes = append(c.writes, append([]byte(nil), p...)) + if c.timeline != nil { + *c.timeline = append(*c.timeline, "write") + } + return c.Conn.Write(p) +} + +type tlsRecord struct { + contentType byte + payload []byte +} + +func parseTLSRecords(t *testing.T, data []byte) []tlsRecord { + t.Helper() + var records []tlsRecord + for len(data) > 0 { + require.GreaterOrEqual(t, len(data), 5, "record header incomplete") + recordLen := int(binary.BigEndian.Uint16(data[3:5])) + require.GreaterOrEqual(t, len(data), 5+recordLen, "record payload truncated") + records = append(records, tlsRecord{ + contentType: data[0], + payload: append([]byte(nil), data[5:5+recordLen]...), + }) + data = data[5+recordLen:] + } + return records +} + +// TestConn_StackedWithRecordFragment mirrors the wrapping order that +// STDClientConfig.Client() produces when record_fragment is enabled: +// tls.Client → tlsspoof.Conn → tf.Conn → raw conn. +// Asserts the decoy is injected and the real handshake arrives split into +// multiple TLS records whose payloads reassemble to the original. +func TestConn_StackedWithRecordFragment(t *testing.T) { + t.Parallel() + payload := decodeClientHello(t) + + client, server := net.Pipe() + defer server.Close() + + fragConn := tf.NewConn(client, context.Background(), false, true, time.Millisecond) spoofer := &fakeSpoofer{} - wrapped := NewConn(client, spoofer, "letsencrypt.org") + wrapped, err := newConn(fragConn, spoofer, "letsencrypt.org") + require.NoError(t, err) - _, err := wrapped.Write([]byte("not a ClientHello")) - require.Error(t, err) - require.Empty(t, spoofer.injected) + serverRead := make(chan []byte, 1) + go func() { serverRead <- readAll(t, server) }() + + _, err = wrapped.Write(payload) + require.NoError(t, err) + require.NoError(t, wrapped.Close()) + forwarded := <-serverRead + + require.Len(t, spoofer.injected, 1, "spoof must inject exactly once") + injected := tf.IndexTLSServerName(spoofer.injected[0]) + require.NotNil(t, injected, "injected payload must parse as ClientHello") + require.Equal(t, "letsencrypt.org", injected.ServerName) + + records := parseTLSRecords(t, forwarded) + require.Greater(t, len(records), 1, "record_fragment must produce multiple records") + var reassembled []byte + for _, r := range records { + require.Equal(t, byte(0x16), r.contentType, "all records must be handshake") + reassembled = append(reassembled, r.payload...) + } + require.Equal(t, payload[5:], reassembled, "record payloads must reassemble to original handshake") +} + +// TestConn_StackedWithPacketFragment is the primary regression test for the +// fragment-only gate fix in STDClientConfig.Client(). It verifies that +// packet-level fragmentation combined with spoof produces: +// - one spoof injection carrying the decoy SNI, +// - multiple separate writes to the underlying conn, +// - an unmodified byte stream when those writes are concatenated +// (no extra record framing). +func TestConn_StackedWithPacketFragment(t *testing.T) { + t.Parallel() + payload := decodeClientHello(t) + + client, server := net.Pipe() + defer server.Close() + + rc := &recordingConn{Conn: client} + fragConn := tf.NewConn(rc, context.Background(), true, false, time.Millisecond) + spoofer := &fakeSpoofer{} + wrapped, err := newConn(fragConn, spoofer, "letsencrypt.org") + require.NoError(t, err) + + serverRead := make(chan []byte, 1) + go func() { serverRead <- readAll(t, server) }() + + _, err = wrapped.Write(payload) + require.NoError(t, err) + require.NoError(t, wrapped.Close()) + forwarded := <-serverRead + + require.Len(t, spoofer.injected, 1, "spoof must inject exactly once") + injected := tf.IndexTLSServerName(spoofer.injected[0]) + require.NotNil(t, injected) + require.Equal(t, "letsencrypt.org", injected.ServerName) + + require.Greater(t, len(rc.writes), 1, "fragment must split the ClientHello into multiple writes") + require.Equal(t, payload, bytes.Join(rc.writes, nil), + "concatenated writes must equal original bytes (no extra framing)") + require.Equal(t, payload, forwarded) +} + +// TestConn_StackedWithBothFragment exercises the combination that produces +// the strongest obfuscation: each chunk becomes its own TLS record and its +// own TCP write. +func TestConn_StackedWithBothFragment(t *testing.T) { + t.Parallel() + payload := decodeClientHello(t) + + client, server := net.Pipe() + defer server.Close() + + rc := &recordingConn{Conn: client} + fragConn := tf.NewConn(rc, context.Background(), true, true, time.Millisecond) + spoofer := &fakeSpoofer{} + wrapped, err := newConn(fragConn, spoofer, "letsencrypt.org") + require.NoError(t, err) + + serverRead := make(chan []byte, 1) + go func() { serverRead <- readAll(t, server) }() + + _, err = wrapped.Write(payload) + require.NoError(t, err) + require.NoError(t, wrapped.Close()) + forwarded := <-serverRead + + require.Len(t, spoofer.injected, 1) + injected := tf.IndexTLSServerName(spoofer.injected[0]) + require.NotNil(t, injected) + require.Equal(t, "letsencrypt.org", injected.ServerName) + + require.Greater(t, len(rc.writes), 1, "split-packet must produce multiple writes") + records := parseTLSRecords(t, forwarded) + require.Greater(t, len(records), 1, "split-record must produce multiple records") + var reassembled []byte + for _, r := range records { + require.Equal(t, byte(0x16), r.contentType) + reassembled = append(reassembled, r.payload...) + } + require.Equal(t, payload[5:], reassembled, + "record payloads must reassemble to the original handshake") +} + +// trackingSpoofer adds the spoof injection to a shared event timeline so +// TestConn_StackedInjectionOrder can prove the decoy precedes the first +// downstream write. +type trackingSpoofer struct { + injected [][]byte + timeline *[]string +} + +func (s *trackingSpoofer) Inject(payload []byte) error { + s.injected = append(s.injected, append([]byte(nil), payload...)) + *s.timeline = append(*s.timeline, "inject") + return nil +} + +func (s *trackingSpoofer) Close() error { return nil } + +// TestConn_StackedInjectionOrder asserts the documented wire order: the +// decoy injection happens before any write reaches the underlying conn. +func TestConn_StackedInjectionOrder(t *testing.T) { + t.Parallel() + payload := decodeClientHello(t) + + client, server := net.Pipe() + defer server.Close() + + var timeline []string + rc := &recordingConn{Conn: client, timeline: &timeline} + fragConn := tf.NewConn(rc, context.Background(), true, true, time.Millisecond) + spoofer := &trackingSpoofer{timeline: &timeline} + wrapped, err := newConn(fragConn, spoofer, "letsencrypt.org") + require.NoError(t, err) + + serverRead := make(chan []byte, 1) + go func() { serverRead <- readAll(t, server) }() + + _, err = wrapped.Write(payload) + require.NoError(t, err) + require.NoError(t, wrapped.Close()) + <-serverRead + + require.NotEmpty(t, timeline) + require.Equal(t, "inject", timeline[0], "decoy must be injected before any downstream write") + require.Contains(t, timeline[1:], "write", "at least one downstream write must follow the inject") } func TestParseMethod(t *testing.T) { diff --git a/common/tlsspoof/integration_test.go b/common/tlsspoof/integration_test.go index b7b07d54b..23a83ff17 100644 --- a/common/tlsspoof/integration_test.go +++ b/common/tlsspoof/integration_test.go @@ -11,7 +11,7 @@ import ( "os" "os/exec" "strings" - "sync/atomic" + "sync" "testing" "time" @@ -21,11 +21,20 @@ import ( func requireRoot(t *testing.T) { t.Helper() if os.Geteuid() != 0 { - t.Fatal("integration test requires root") + t.Skip("integration test requires root; re-run with `go test -exec sudo`") } } func tcpdumpObserver(t *testing.T, iface string, port uint16, needle string, do func(), wait time.Duration) bool { + t.Helper() + return tcpdumpObserverMulti(t, iface, port, []string{needle}, do, wait)[needle] +} + +// tcpdumpObserverMulti captures tcpdump output while do() executes and reports +// which of the provided needles were observed in the raw ASCII dump. Use this +// to assert that distinct payloads (e.g. fake vs real ClientHello) are both on +// the wire. +func tcpdumpObserverMulti(t *testing.T, iface string, port uint16, needles []string, do func(), wait time.Duration) map[string]bool { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), wait) defer cancel() @@ -62,16 +71,22 @@ func tcpdumpObserver(t *testing.T, iface string, port uint16, needle string, do t.Fatal("tcpdump did not attach within 2s") } - var found atomic.Bool + var access sync.Mutex + found := make(map[string]bool, len(needles)) readerDone := make(chan struct{}) go func() { defer close(readerDone) scanner := bufio.NewScanner(stdout) scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) for scanner.Scan() { - if strings.Contains(scanner.Text(), needle) { - found.Store(true) + line := scanner.Text() + access.Lock() + for _, needle := range needles { + if !found[needle] && strings.Contains(line, needle) { + found[needle] = true + } } + access.Unlock() } }() @@ -80,7 +95,13 @@ func tcpdumpObserver(t *testing.T, iface string, port uint16, needle string, do time.Sleep(200 * time.Millisecond) _ = cmd.Process.Signal(os.Interrupt) <-readerDone - return found.Load() + access.Lock() + defer access.Unlock() + result := make(map[string]bool, len(needles)) + for _, needle := range needles { + result[needle] = found[needle] + } + return result } func dialLocalEchoServer(t *testing.T) (client net.Conn, serverPort uint16) { diff --git a/common/tlsspoof/integration_tls_test.go b/common/tlsspoof/integration_tls_test.go new file mode 100644 index 000000000..d179c3841 --- /dev/null +++ b/common/tlsspoof/integration_tls_test.go @@ -0,0 +1,118 @@ +//go:build linux || darwin + +package tlsspoof + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "io" + "math/big" + "net" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// generateSelfSignedCert returns a TLS certificate valid for the given SAN. +func generateSelfSignedCert(t *testing.T, commonName string, sans ...string) tls.Certificate { + t.Helper() + priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128)) + require.NoError(t, err) + template := x509.Certificate{ + SerialNumber: serial, + Subject: pkix.Name{CommonName: commonName}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + DNSNames: sans, + } + der, err := x509.CreateCertificate(rand.Reader, &template, &template, &priv.PublicKey, priv) + require.NoError(t, err) + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}) + keyDER, err := x509.MarshalPKCS8PrivateKey(priv) + require.NoError(t, err) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyDER}) + cert, err := tls.X509KeyPair(certPEM, keyPEM) + require.NoError(t, err) + return cert +} + +// TestIntegrationConn_RealTLSHandshake drives a real crypto/tls ClientHello +// through the spoofer and asserts the on-wire fake packet carries the fake SNI +// while the server receives the real SNI. This exercises the full +// `tls.Client(wrapped, config).Handshake()` path rather than a static hex +// payload, matching what user-facing code hits. +func TestIntegrationConn_RealTLSHandshake(t *testing.T) { + requireRoot(t) + const realSNI = "real.test" + const fakeSNI = "fake.test" + + serverCert := generateSelfSignedCert(t, realSNI, realSNI) + tlsConfig := &tls.Config{Certificates: []tls.Certificate{serverCert}} + + listener, err := tls.Listen("tcp4", "127.0.0.1:0", tlsConfig) + require.NoError(t, err) + t.Cleanup(func() { listener.Close() }) + + serverSNI := make(chan string, 1) + go func() { + conn, err := listener.Accept() + if err != nil { + return + } + defer conn.Close() + tlsConn := conn.(*tls.Conn) + _ = tlsConn.SetDeadline(time.Now().Add(3 * time.Second)) + if handshakeErr := tlsConn.Handshake(); handshakeErr != nil { + serverSNI <- "handshake-error:" + handshakeErr.Error() + return + } + serverSNI <- tlsConn.ConnectionState().ServerName + _, _ = io.Copy(io.Discard, conn) + }() + + addr := listener.Addr().(*net.TCPAddr) + serverPort := uint16(addr.Port) + raw, err := net.Dial("tcp4", addr.String()) + require.NoError(t, err) + t.Cleanup(func() { raw.Close() }) + + wrapped, err := NewConn(raw, MethodWrongSequence, fakeSNI) + require.NoError(t, err) + + clientConfig := &tls.Config{ + ServerName: realSNI, + InsecureSkipVerify: true, + } + tlsClient := tls.Client(wrapped, clientConfig) + t.Cleanup(func() { tlsClient.Close() }) + + seen := tcpdumpObserverMulti(t, loopbackInterface, serverPort, + []string{realSNI, fakeSNI}, func() { + _ = tlsClient.SetDeadline(time.Now().Add(3 * time.Second)) + err := tlsClient.Handshake() + require.NoError(t, err, "TLS handshake must succeed (wrong-sequence fake is dropped by peer)") + }, 4*time.Second) + + require.True(t, seen[realSNI], + "real ClientHello on the wire must contain original SNI %q", realSNI) + require.True(t, seen[fakeSNI], + "fake ClientHello on the wire must contain fake SNI %q", fakeSNI) + + select { + case sniOnServer := <-serverSNI: + require.Equal(t, realSNI, sniOnServer, + "TLS server must see the real SNI (fake packet dropped by peer TCP stack)") + case <-time.After(3 * time.Second): + t.Fatal("TLS server did not complete handshake") + } +} diff --git a/common/tlsspoof/integration_unix_test.go b/common/tlsspoof/integration_unix_test.go index 9ec5760c7..0f4585fd8 100644 --- a/common/tlsspoof/integration_unix_test.go +++ b/common/tlsspoof/integration_unix_test.go @@ -15,13 +15,11 @@ import ( func TestIntegrationSpoofer_WrongChecksum(t *testing.T) { requireRoot(t) client, serverPort := dialLocalEchoServer(t) - spoofer, err := NewSpoofer(client, MethodWrongChecksum) + spoofer, err := newRawSpoofer(client, MethodWrongChecksum) require.NoError(t, err) defer spoofer.Close() - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) - fake, err := rewriteSNI(payload, "letsencrypt.org") + fake, err := buildFakeClientHello("letsencrypt.org") require.NoError(t, err) captured := tcpdumpObserver(t, loopbackInterface, serverPort, "letsencrypt.org", func() { @@ -33,13 +31,11 @@ func TestIntegrationSpoofer_WrongChecksum(t *testing.T) { func TestIntegrationSpoofer_WrongSequence(t *testing.T) { requireRoot(t) client, serverPort := dialLocalEchoServer(t) - spoofer, err := NewSpoofer(client, MethodWrongSequence) + spoofer, err := newRawSpoofer(client, MethodWrongSequence) require.NoError(t, err) defer spoofer.Close() - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) - fake, err := rewriteSNI(payload, "letsencrypt.org") + fake, err := buildFakeClientHello("letsencrypt.org") require.NoError(t, err) captured := tcpdumpObserver(t, loopbackInterface, serverPort, "letsencrypt.org", func() { @@ -51,13 +47,11 @@ func TestIntegrationSpoofer_WrongSequence(t *testing.T) { func TestIntegrationSpoofer_IPv6_WrongChecksum(t *testing.T) { requireRoot(t) client, serverPort := dialLocalEchoServerIPv6(t) - spoofer, err := NewSpoofer(client, MethodWrongChecksum) + spoofer, err := newRawSpoofer(client, MethodWrongChecksum) require.NoError(t, err) defer spoofer.Close() - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) - fake, err := rewriteSNI(payload, "letsencrypt.org") + fake, err := buildFakeClientHello("letsencrypt.org") require.NoError(t, err) captured := tcpdumpObserver(t, loopbackInterface, serverPort, "letsencrypt.org", func() { @@ -69,13 +63,11 @@ func TestIntegrationSpoofer_IPv6_WrongChecksum(t *testing.T) { func TestIntegrationSpoofer_IPv6_WrongSequence(t *testing.T) { requireRoot(t) client, serverPort := dialLocalEchoServerIPv6(t) - spoofer, err := NewSpoofer(client, MethodWrongSequence) + spoofer, err := newRawSpoofer(client, MethodWrongSequence) require.NoError(t, err) defer spoofer.Close() - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) - fake, err := rewriteSNI(payload, "letsencrypt.org") + fake, err := buildFakeClientHello("letsencrypt.org") require.NoError(t, err) captured := tcpdumpObserver(t, loopbackInterface, serverPort, "letsencrypt.org", func() { @@ -95,6 +87,76 @@ func TestIntegrationConn_IPv6_InjectsThenForwardsRealCH(t *testing.T) { runInjectsThenForwardsRealCH(t, "tcp6", "[::1]:0") } +// TestIntegrationConn_FakeAndRealHaveDistinctSNIs asserts that the on-wire fake +// packet carries the fake SNI (letsencrypt.org) AND the real packet still +// carries the original SNI (github.com). If the builder regresses to producing +// empty or mismatched bytes, the fake-SNI needle will be missing. +func TestIntegrationConn_FakeAndRealHaveDistinctSNIs(t *testing.T) { + requireRoot(t) + runFakeAndRealHaveDistinctSNIs(t, "tcp4", "127.0.0.1:0", "letsencrypt.org") +} + +func TestIntegrationConn_IPv6_FakeAndRealHaveDistinctSNIs(t *testing.T) { + requireRoot(t) + runFakeAndRealHaveDistinctSNIs(t, "tcp6", "[::1]:0", "letsencrypt.org") +} + +func runFakeAndRealHaveDistinctSNIs(t *testing.T, network, address, fakeSNI string) { + t.Helper() + const originalSNI = "github.com" + require.NotEqual(t, originalSNI, fakeSNI) + + listener, err := net.Listen(network, address) + require.NoError(t, err) + + serverReceived := make(chan []byte, 1) + go func() { + conn, err := listener.Accept() + if err != nil { + return + } + defer conn.Close() + _ = conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + got, _ := io.ReadAll(conn) + serverReceived <- got + }() + + addr := listener.Addr().(*net.TCPAddr) + serverPort := uint16(addr.Port) + client, err := net.Dial(network, addr.String()) + require.NoError(t, err) + t.Cleanup(func() { + client.Close() + listener.Close() + }) + + wrapped, err := NewConn(client, MethodWrongSequence, fakeSNI) + require.NoError(t, err) + + payload, err := hex.DecodeString(realClientHello) + require.NoError(t, err) + + seen := tcpdumpObserverMulti(t, loopbackInterface, serverPort, + []string{originalSNI, fakeSNI}, func() { + n, err := wrapped.Write(payload) + require.NoError(t, err) + require.Equal(t, len(payload), n) + }, 3*time.Second) + require.True(t, seen[originalSNI], + "real ClientHello must carry original SNI %q on the wire", originalSNI) + require.True(t, seen[fakeSNI], + "fake ClientHello must carry fake SNI %q on the wire", fakeSNI) + + _ = wrapped.Close() + select { + case got := <-serverReceived: + require.Equal(t, payload, got, + "server must receive real ClientHello unchanged (wrong-sequence fake must be dropped)") + case <-time.After(2 * time.Second): + t.Fatal("echo server did not receive real ClientHello") + } +} + func runInjectsThenForwardsRealCH(t *testing.T, network, address string) { t.Helper() listener, err := net.Listen(network, address) @@ -121,9 +183,8 @@ func runInjectsThenForwardsRealCH(t *testing.T, network, address string) { listener.Close() }) - spoofer, err := NewSpoofer(client, MethodWrongSequence) + wrapped, err := NewConn(client, MethodWrongSequence, "letsencrypt.org") require.NoError(t, err) - wrapped := NewConn(client, spoofer, "letsencrypt.org") payload, err := hex.DecodeString(realClientHello) require.NoError(t, err) diff --git a/common/tlsspoof/integration_windows_test.go b/common/tlsspoof/integration_windows_test.go index d3f823841..b0461a31b 100644 --- a/common/tlsspoof/integration_windows_test.go +++ b/common/tlsspoof/integration_windows_test.go @@ -12,11 +12,11 @@ import ( "github.com/stretchr/testify/require" ) -func newSpoofer(t *testing.T, conn net.Conn, method Method) Spoofer { +func newSpoofer(t *testing.T, conn net.Conn, method Method) rawSpoofer { t.Helper() - spoofer, err := NewSpoofer(conn, method) + s, err := newRawSpoofer(conn, method) require.NoError(t, err) - return spoofer + return s } // Basic lifecycle: opening a spoofer against a live TCP conn installs @@ -46,11 +46,10 @@ func TestIntegrationSpooferOpenClose(t *testing.T) { require.NoError(t, spoofer.Close()) } -// End-to-end: Conn.Write injects a fake ClientHello with a rewritten -// SNI, then forwards the real ClientHello. With wrong-sequence, the -// fake lands before the connection's send-next sequence — the peer TCP -// stack treats it as already-received and only surfaces the real bytes -// to the echo server. +// End-to-end: Conn.Write injects a fake ClientHello with a fresh SNI, then +// forwards the real ClientHello. With wrong-sequence, the fake lands before +// the connection's send-next sequence — the peer TCP stack treats it as +// already-received and only surfaces the real bytes to the echo server. func TestIntegrationConnInjectsThenForwardsRealCH(t *testing.T) { listener, err := net.Listen("tcp4", "127.0.0.1:0") require.NoError(t, err) @@ -72,8 +71,8 @@ func TestIntegrationConnInjectsThenForwardsRealCH(t *testing.T) { require.NoError(t, err) t.Cleanup(func() { client.Close() }) - spoofer := newSpoofer(t, client, MethodWrongSequence) - wrapped := NewConn(client, spoofer, "letsencrypt.org") + wrapped, err := NewConn(client, MethodWrongSequence, "letsencrypt.org") + require.NoError(t, err) payload, err := hex.DecodeString(realClientHello) require.NoError(t, err) @@ -94,7 +93,7 @@ func TestIntegrationConnInjectsThenForwardsRealCH(t *testing.T) { // Inject before any kernel payload: stages the fake, then Write flushes // the real CH. Same terminal expectation as the Conn variant but via the -// Spoofer primitive directly. +// raw spoofer primitive directly. func TestIntegrationSpooferInjectThenWrite(t *testing.T) { listener, err := net.Listen("tcp4", "127.0.0.1:0") require.NoError(t, err) @@ -119,12 +118,12 @@ func TestIntegrationSpooferInjectThenWrite(t *testing.T) { spoofer := newSpoofer(t, client, MethodWrongSequence) t.Cleanup(func() { spoofer.Close() }) - payload, err := hex.DecodeString(realClientHello) - require.NoError(t, err) - fake, err := rewriteSNI(payload, "letsencrypt.org") + fake, err := buildFakeClientHello("letsencrypt.org") require.NoError(t, err) require.NoError(t, spoofer.Inject(fake)) + payload, err := hex.DecodeString(realClientHello) + require.NoError(t, err) n, err := client.Write(payload) require.NoError(t, err) require.Equal(t, len(payload), n) diff --git a/common/tlsspoof/packet_test.go b/common/tlsspoof/packet_test.go index 992a96840..5c6d5b6be 100644 --- a/common/tlsspoof/packet_test.go +++ b/common/tlsspoof/packet_test.go @@ -75,3 +75,62 @@ func TestBuildTCPSegment_MixedFamilyPanics(t *testing.T) { buildTCPSegment(src, dst, 0, 0, nil, false) }) } + +func TestBuildSpoofFrame_WrongSequence(t *testing.T) { + t.Parallel() + src := netip.MustParseAddrPort("10.0.0.1:54321") + dst := netip.MustParseAddrPort("1.2.3.4:443") + payload := []byte("fake-client-hello") + const sendNext uint32 = 10_000 + frame, err := buildSpoofFrame(MethodWrongSequence, src, dst, sendNext, 20_000, payload) + require.NoError(t, err) + + tcp := header.TCP(frame[header.IPv4MinimumSize:]) + require.Equal(t, sendNext-uint32(len(payload)), tcp.SequenceNumber(), + "wrong-sequence places the fake at sendNext-len(payload)") + require.True(t, tcp.Flags().Contains(header.TCPFlagAck|header.TCPFlagPsh)) + + // Checksum must still be valid — only the sequence number is wrong. + payloadChecksum := checksum.Checksum(payload, 0) + require.True(t, tcp.IsChecksumValid( + tcpip.AddrFrom4(src.Addr().As4()), + tcpip.AddrFrom4(dst.Addr().As4()), + payloadChecksum, + uint16(len(payload)), + )) +} + +func TestBuildSpoofFrame_WrongChecksum(t *testing.T) { + t.Parallel() + src := netip.MustParseAddrPort("10.0.0.1:54321") + dst := netip.MustParseAddrPort("1.2.3.4:443") + payload := []byte("fake-client-hello") + const sendNext uint32 = 5_000 + frame, err := buildSpoofFrame(MethodWrongChecksum, src, dst, sendNext, 20_000, payload) + require.NoError(t, err) + + tcp := header.TCP(frame[header.IPv4MinimumSize:]) + require.Equal(t, sendNext, tcp.SequenceNumber(), + "wrong-checksum keeps the real sequence number") + + payloadChecksum := checksum.Checksum(payload, 0) + require.False(t, tcp.IsChecksumValid( + tcpip.AddrFrom4(src.Addr().As4()), + tcpip.AddrFrom4(dst.Addr().As4()), + payloadChecksum, + uint16(len(payload)), + )) + require.True(t, header.IPv4(frame[:header.IPv4MinimumSize]).IsChecksumValid(), + "IPv4 checksum must remain valid so the router forwards the packet") +} + +func TestBuildSpoofTCPSegment_EncodesWithoutIPHeader(t *testing.T) { + t.Parallel() + src := netip.MustParseAddrPort("[fe80::1]:54321") + dst := netip.MustParseAddrPort("[2606:4700::1]:443") + payload := []byte("fake-client-hello") + segment, err := buildSpoofTCPSegment(MethodWrongSequence, src, dst, 1000, 2000, payload) + require.NoError(t, err) + require.Equal(t, tcpHeaderLen+len(payload), len(segment), + "segment must be TCP header + payload, no IP header") +} diff --git a/common/tlsspoof/raw_darwin.go b/common/tlsspoof/raw_darwin.go index 99c9a5c66..ab3168769 100644 --- a/common/tlsspoof/raw_darwin.go +++ b/common/tlsspoof/raw_darwin.go @@ -9,6 +9,7 @@ import ( "sync" "syscall" + "github.com/sagernet/sing-tun/gtcpip/header" E "github.com/sagernet/sing/common/exceptions" "golang.org/x/sys/unix" @@ -34,14 +35,26 @@ const ( darwinXtcpcbRcvNxtOffset = 80 ) -var darwinStructSize = sync.OnceValue(func() int { - value, _ := syscall.Sysctl("kern.osrelease") - major, _, _ := strings.Cut(value, ".") - n, _ := strconv.ParseInt(major, 10, 64) - if n >= 22 { - return 408 +// darwinStructSize returns the size of xinpcb_n for the running Darwin kernel. +// Darwin 22 (macOS 13 Ventura) grew the struct from 384 to 408 bytes; there is +// no ABI-stable way to read it, so we key off the kernel version. +var darwinStructSize = sync.OnceValues(func() (int, error) { + value, err := syscall.Sysctl("kern.osrelease") + if err != nil { + return 0, E.Cause(err, "sysctl kern.osrelease") } - return 384 + major, _, ok := strings.Cut(value, ".") + if !ok { + return 0, E.New("unexpected kern.osrelease format: ", value) + } + n, err := strconv.ParseInt(major, 10, 64) + if err != nil { + return 0, E.Cause(err, "parse kern.osrelease major version: ", value) + } + if n >= 22 { + return 408, nil + } + return 384, nil }) type darwinSpoofer struct { @@ -54,7 +67,7 @@ type darwinSpoofer struct { receiveNext uint32 } -func newRawSpoofer(conn net.Conn, method Method) (Spoofer, error) { +func newRawSpoofer(conn net.Conn, method Method) (rawSpoofer, error) { _, src, dst, err := tcpEndpoints(conn) if err != nil { return nil, err @@ -87,7 +100,10 @@ func readDarwinTCPSequence(src, dst netip.AddrPort) (uint32, uint32, error) { if err != nil { return 0, 0, E.Cause(err, "sysctl net.inet.tcp.pcblist_n") } - structSize := darwinStructSize() + structSize, err := darwinStructSize() + if err != nil { + return 0, 0, err + } itemSize := structSize + darwinTCPExtraSize for i := darwinXinpgenSize; i+itemSize <= len(buffer); i += itemSize { inpcb := buffer[i : i+darwinXsocketOffset] @@ -160,10 +176,9 @@ func (s *darwinSpoofer) Inject(payload []byte) error { // Darwin inherits the historical BSD quirk: with IP_HDRINCL the kernel // expects ip_len and ip_off in host byte order, not network byte order. // Apple's rip_output swaps them back before transmission. - totalLen := binary.BigEndian.Uint16(frame[2:4]) - binary.NativeEndian.PutUint16(frame[2:4], totalLen) - fragOff := binary.BigEndian.Uint16(frame[6:8]) - binary.NativeEndian.PutUint16(frame[6:8], fragOff) + ip := header.IPv4(frame) + ip.SetTotalLengthDarwinRaw(ip.TotalLength()) + ip.SetFlagsFragmentOffsetDarwinRaw(ip.Flags(), ip.FragmentOffset()) err = unix.Sendto(s.rawFD, frame, 0, s.rawSockAddr) if err != nil { return E.Cause(err, "sendto raw socket") diff --git a/common/tlsspoof/raw_linux.go b/common/tlsspoof/raw_linux.go index cb694aba9..f82fbc9ef 100644 --- a/common/tlsspoof/raw_linux.go +++ b/common/tlsspoof/raw_linux.go @@ -29,7 +29,7 @@ type linuxSpoofer struct { receiveNext uint32 } -func newRawSpoofer(conn net.Conn, method Method) (Spoofer, error) { +func newRawSpoofer(conn net.Conn, method Method) (rawSpoofer, error) { tcpConn, src, dst, err := tcpEndpoints(conn) if err != nil { return nil, err @@ -66,22 +66,34 @@ func openLinuxRawSocket(dst netip.AddrPort) (int, unix.Sockaddr, error) { unix.Close(fd) return -1, nil, E.Cause(err, "set IPV6_HDRINCL") } - sockaddr := &unix.SockaddrInet6{Port: int(dst.Port())} - sockaddr.Addr = dst.Addr().As16() + // Linux raw IPv6 sockets interpret sin6_port as a nexthdr protocol number + // (see raw(7)); any value other than 0 or the socket's IPPROTO_TCP causes + // sendto to fail with EINVAL. The destination is already encoded in the + // user-supplied IPv6 header under IPV6_HDRINCL. + sockaddr := &unix.SockaddrInet6{Addr: dst.Addr().As16()} return fd, sockaddr, nil } // loadSequenceNumbers puts the socket briefly into TCP_REPAIR mode to read // snd_nxt and rcv_nxt from the kernel. TCP_REPAIR requires CAP_NET_ADMIN; // callers must run as root or grant both CAP_NET_RAW and CAP_NET_ADMIN. +// +// If the TCP_REPAIR_OFF revert fails, the socket would stay in TCP_REPAIR +// state and subsequent Write() calls would silently buffer instead of sending. +// Surface that error so callers can abort. func (s *linuxSpoofer) loadSequenceNumbers(tcpConn *net.TCPConn) error { - return control.Conn(tcpConn, func(raw uintptr) error { + return control.Conn(tcpConn, func(raw uintptr) (err error) { fd := int(raw) - err := unix.SetsockoptInt(fd, unix.IPPROTO_TCP, unix.TCP_REPAIR, unix.TCP_REPAIR_ON) + err = unix.SetsockoptInt(fd, unix.IPPROTO_TCP, unix.TCP_REPAIR, unix.TCP_REPAIR_ON) if err != nil { return E.Cause(err, "enter TCP_REPAIR (need CAP_NET_ADMIN)") } - defer unix.SetsockoptInt(fd, unix.IPPROTO_TCP, unix.TCP_REPAIR, unix.TCP_REPAIR_OFF) + defer func() { + offErr := unix.SetsockoptInt(fd, unix.IPPROTO_TCP, unix.TCP_REPAIR, unix.TCP_REPAIR_OFF) + if err == nil && offErr != nil { + err = E.Cause(offErr, "leave TCP_REPAIR") + } + }() err = unix.SetsockoptInt(fd, unix.IPPROTO_TCP, unix.TCP_REPAIR_QUEUE, tcpSendQueue) if err != nil { diff --git a/common/tlsspoof/raw_stub.go b/common/tlsspoof/raw_stub.go index a2da87d6b..7edf2441a 100644 --- a/common/tlsspoof/raw_stub.go +++ b/common/tlsspoof/raw_stub.go @@ -10,6 +10,6 @@ import ( const PlatformSupported = false -func newRawSpoofer(conn net.Conn, method Method) (Spoofer, error) { +func newRawSpoofer(conn net.Conn, method Method) (rawSpoofer, error) { return nil, E.New("tls_spoof: unsupported platform") } diff --git a/common/tlsspoof/raw_windows.go b/common/tlsspoof/raw_windows.go index b6961169f..9f6553f1b 100644 --- a/common/tlsspoof/raw_windows.go +++ b/common/tlsspoof/raw_windows.go @@ -25,11 +25,15 @@ const PlatformSupported = true // bounds the pathological case where the kernel buffers the packet. const closeGracePeriod = 2 * time.Second +// windowsSpoofer uses a single WinDivert handle for both capture and +// injection. Sequential Send() calls on one handle traverse one driver queue, +// so the fake provably precedes the released real on the wire — a guarantee +// two separate handles cannot make because cross-handle order depends on the +// scheduler. type windowsSpoofer struct { method Method src, dst netip.AddrPort divertH *windivert.Handle - injectH *windivert.Handle fakeReady chan []byte // buffered(1): staged by Inject done chan struct{} // closed by run() on exit @@ -37,12 +41,11 @@ type windowsSpoofer struct { runErr atomic.Pointer[error] } -func newRawSpoofer(conn net.Conn, method Method) (Spoofer, error) { +func newRawSpoofer(conn net.Conn, method Method) (rawSpoofer, error) { _, src, dst, err := tcpEndpoints(conn) if err != nil { return nil, err } - filter, err := windivert.OutboundTCP(src, dst) if err != nil { return nil, err @@ -51,17 +54,11 @@ func newRawSpoofer(conn net.Conn, method Method) (Spoofer, error) { if err != nil { return nil, E.Cause(err, "tls_spoof: open WinDivert") } - injectH, err := windivert.Open(nil, windivert.LayerNetwork, 0, windivert.FlagSendOnly) - if err != nil { - divertH.Close() - return nil, E.Cause(err, "tls_spoof: open WinDivert") - } s := &windowsSpoofer{ method: method, src: src, dst: dst, divertH: divertH, - injectH: injectH, fakeReady: make(chan []byte, 1), done: make(chan struct{}), } @@ -91,7 +88,6 @@ func (s *windowsSpoofer) Close() error { s.divertH.Close() <-s.done } - s.injectH.Close() }) if p := s.runErr.Load(); p != nil { return *p @@ -119,9 +115,17 @@ func (s *windowsSpoofer) run() { pkt := buf[:n] seq, ack, payloadLen, ok := parseTCPFields(pkt, addr.IPv6()) if !ok { - // Malformed / not TCP — shouldn't match our filter, but be safe. - _, _ = s.divertH.Send(pkt, &addr) - continue + // Our filter is OutboundTCP(src, dst); a non-TCP or truncated + // match means driver state is suspect. Re-inject so the kernel + // still sees the byte stream, then abort — continuing would risk + // reordering against an unknown reference point. + _, sendErr := s.divertH.Send(pkt, &addr) + if sendErr != nil { + s.recordErr(E.Cause(sendErr, "windivert re-inject malformed")) + return + } + s.recordErr(E.New("windivert received malformed packet matching spoof filter")) + return } if payloadLen == 0 { // Handshake ACK, keepalive, FIN — pass through unchanged. @@ -159,7 +163,7 @@ func (s *windowsSpoofer) run() { // Force both to 1 to keep our bytes intact. fakeAddr.SetIPChecksum(true) fakeAddr.SetTCPChecksum(true) - _, err = s.injectH.Send(frame, &fakeAddr) + _, err = s.divertH.Send(frame, &fakeAddr) if err != nil { s.recordErr(E.Cause(err, "windivert inject fake")) return diff --git a/common/tlsspoof/spoof.go b/common/tlsspoof/spoof.go index 2a27ec328..1bca5693f 100644 --- a/common/tlsspoof/spoof.go +++ b/common/tlsspoof/spoof.go @@ -40,40 +40,54 @@ func (m Method) String() string { } } -type Spoofer interface { +type rawSpoofer interface { Inject(payload []byte) error Close() error } -func NewSpoofer(conn net.Conn, method Method) (Spoofer, error) { - return newRawSpoofer(conn, method) -} - type Conn struct { net.Conn - spoofer Spoofer - fakeSNI string - injected bool + spoofer rawSpoofer + fakeHello []byte + injected bool } -func NewConn(conn net.Conn, spoofer Spoofer, fakeSNI string) *Conn { - return &Conn{ - Conn: conn, - spoofer: spoofer, - fakeSNI: fakeSNI, +func NewConn(conn net.Conn, method Method, fakeSNI string) (*Conn, error) { + spoofer, err := newRawSpoofer(conn, method) + if err != nil { + return nil, err } + result, err := newConn(conn, spoofer, fakeSNI) + if err != nil { + spoofer.Close() + return nil, err + } + return result, nil } -func (c *Conn) Write(b []byte) (int, error) { +func newConn(conn net.Conn, spoofer rawSpoofer, fakeSNI string) (*Conn, error) { + fakeHello, err := buildFakeClientHello(fakeSNI) + if err != nil { + return nil, E.Cause(err, "tls_spoof: build fake ClientHello") + } + return &Conn{ + Conn: conn, + spoofer: spoofer, + fakeHello: fakeHello, + }, nil +} + +func (c *Conn) Write(b []byte) (n int, err error) { if c.injected { return c.Conn.Write(b) } - defer c.spoofer.Close() - fake, err := rewriteSNI(b, c.fakeSNI) - if err != nil { - return 0, E.Cause(err, "tls_spoof: rewrite SNI") - } - err = c.spoofer.Inject(fake) + defer func() { + closeErr := c.spoofer.Close() + if err == nil && closeErr != nil { + err = E.Cause(closeErr, "tls_spoof: close spoofer") + } + }() + err = c.spoofer.Inject(c.fakeHello) if err != nil { return 0, E.Cause(err, "tls_spoof: inject") } @@ -83,7 +97,7 @@ func (c *Conn) Write(b []byte) (int, error) { func (c *Conn) Close() error { return E.Append(c.Conn.Close(), c.spoofer.Close(), func(e error) error { - return E.Cause(e, "close spoofer") + return E.Cause(e, "tls_spoof: close spoofer") }) } diff --git a/common/windivert/handle_windows.go b/common/windivert/handle_windows.go index e7f5ae673..1d7aebfde 100644 --- a/common/windivert/handle_windows.go +++ b/common/windivert/handle_windows.go @@ -110,9 +110,13 @@ func validateOpenArgs(layer Layer, priority int16, flags Flag) error { if priority < PriorityLowest || priority > PriorityHighest { return E.New("windivert: priority out of range") } - if flags&^FlagSendOnly != 0 { + const supportedFlags = FlagSniff | FlagSendOnly + if flags&^supportedFlags != 0 { return E.New("windivert: unknown flag bits") } + if flags&FlagSniff != 0 && flags&FlagSendOnly != 0 { + return E.New("windivert: FlagSniff and FlagSendOnly are mutually exclusive") + } return nil } diff --git a/common/windivert/handle_windows_test.go b/common/windivert/handle_windows_test.go index dd05ce7b0..73dfbb166 100644 --- a/common/windivert/handle_windows_test.go +++ b/common/windivert/handle_windows_test.go @@ -100,6 +100,9 @@ func TestValidateOpenArgsFlags(t *testing.T) { t.Parallel() require.NoError(t, validateOpenArgs(LayerNetwork, 0, 0)) require.NoError(t, validateOpenArgs(LayerNetwork, 0, FlagSendOnly)) + require.NoError(t, validateOpenArgs(LayerNetwork, 0, FlagSniff)) + // Sniff and send-only describe contradictory handle roles. + require.Error(t, validateOpenArgs(LayerNetwork, 0, FlagSniff|FlagSendOnly)) // Unknown flag bits must be rejected to surface caller mistakes early. require.Error(t, validateOpenArgs(LayerNetwork, 0, Flag(0x10))) require.Error(t, validateOpenArgs(LayerNetwork, 0, FlagSendOnly|Flag(0x10))) diff --git a/common/windivert/windivert.go b/common/windivert/windivert.go index e9a8fc954..9d309886c 100644 --- a/common/windivert/windivert.go +++ b/common/windivert/windivert.go @@ -23,7 +23,14 @@ const LayerNetwork Layer = 0 type Flag uint64 -const FlagSendOnly Flag = 0x0008 +const ( + // FlagSniff opens a passive observer: the driver copies matching packets + // to userspace without removing them from the network stack. Send is not + // required (and not allowed) on a sniffing handle. + FlagSniff Flag = 0x0001 + // FlagSendOnly opens a write-only injection handle; Recv is not allowed. + FlagSendOnly Flag = 0x0008 +) const ( PriorityHighest int16 = 30000