diff --git a/common/net/upstream.go b/common/net/upstream.go new file mode 100644 index 00000000..8ae47454 --- /dev/null +++ b/common/net/upstream.go @@ -0,0 +1,28 @@ +package net + +import "net" + +type netConn interface { + NetConn() net.Conn +} + +// FindUpstream finds a value in an upstream wrapper chain. If accept rejects a +// matching value, the search continues so an outer wrapper cannot hide a valid +// inner value of the same type. +func FindUpstream[T any](value any, accept func(T) bool) (T, bool) { + for value != nil { + if candidate, ok := value.(T); ok && (accept == nil || accept(candidate)) { + return candidate, true + } + switch wrapper := value.(type) { + case WithUpstream: + value = wrapper.Upstream() + case netConn: + value = wrapper.NetConn() + default: + value = nil + } + } + var zero T + return zero, false +} diff --git a/transport/jls/jls.go b/transport/jls/jls.go index 0ce0f843..fd3b9e67 100644 --- a/transport/jls/jls.go +++ b/transport/jls/jls.go @@ -195,14 +195,14 @@ func Server(ctx context.Context, conn net.Conn, config *ServerConfig) (net.Conn, } func UserFromConn(conn net.Conn) (string, bool) { - tlsConn, ok := conn.(*tls.Conn) + tlsConn, ok := N.FindUpstream(conn, func(tlsConn *tls.Conn) bool { + state := tlsConn.ConnectionState().JLS + return state.Authenticated && state.User != "" + }) if !ok { return "", false } state := tlsConn.ConnectionState().JLS - if !state.Authenticated || state.User == "" { - return "", false - } return state.User, true } diff --git a/transport/jls/jls_test.go b/transport/jls/jls_test.go index 3e12e06c..9fe7e97c 100644 --- a/transport/jls/jls_test.go +++ b/transport/jls/jls_test.go @@ -9,6 +9,7 @@ import ( "testing" "time" + N "github.com/metacubex/mihomo/common/net" "github.com/metacubex/mihomo/component/ca" "github.com/metacubex/http" @@ -56,7 +57,8 @@ func testJLSClientServer(t *testing.T, clientFingerprint string) { serverDone <- errors.New("server did not authenticate JLS user") return } - if authenticatedUser, ok := UserFromConn(conn); !ok || authenticatedUser != user.Username { + outerTLS := tls.Server(N.NewBufferedConn(N.NewCachedConn(conn, nil)), &tls.Config{}) + if authenticatedUser, ok := UserFromConn(outerTLS); !ok || authenticatedUser != user.Username { serverDone <- errors.New("server did not expose JLS user") return } diff --git a/transport/shadowtls/server.go b/transport/shadowtls/server.go index d5d68820..61e0c98c 100644 --- a/transport/shadowtls/server.go +++ b/transport/shadowtls/server.go @@ -289,7 +289,7 @@ type authenticatedConn struct { func (c *authenticatedConn) Upstream() any { return c.Conn } func UserFromConn(conn net.Conn) (string, bool) { - authenticated, ok := conn.(*authenticatedConn) + authenticated, ok := N.FindUpstream[*authenticatedConn](conn, nil) if !ok { return "", false } diff --git a/transport/shadowtls/server_test.go b/transport/shadowtls/server_test.go index ce54f4f7..28c1471d 100644 --- a/transport/shadowtls/server_test.go +++ b/transport/shadowtls/server_test.go @@ -13,6 +13,7 @@ import ( "testing" "time" + N "github.com/metacubex/mihomo/common/net" "github.com/metacubex/mihomo/component/ca" "github.com/metacubex/tls" ) @@ -98,6 +99,21 @@ func TestServer(t *testing.T) { } } +func TestUserFromConnThroughWrappers(t *testing.T) { + left, right := net.Pipe() + t.Cleanup(func() { + _ = left.Close() + _ = right.Close() + }) + authenticated := &authenticatedConn{Conn: left, user: "test-user"} + conn := N.NewBufferedConn(N.NewCachedConn(authenticated, nil)) + + user, ok := UserFromConn(conn) + if !ok || user != authenticated.user { + t.Fatalf("UserFromConn() = (%q, %v), want (%q, true)", user, ok, authenticated.user) + } +} + func TestV3UnauthenticatedConnectionFallsBack(t *testing.T) { camouflageAddr := startCamouflageServer(t, false) serverConfig := newTestServerConfig(t, 3, camouflageAddr) diff --git a/transport/simple-obfs/http_server.go b/transport/simple-obfs/http_server.go index c1925acf..10265f68 100644 --- a/transport/simple-obfs/http_server.go +++ b/transport/simple-obfs/http_server.go @@ -22,6 +22,10 @@ type HTTPObfsServer struct { firstResponse bool } +func (hos *HTTPObfsServer) Upstream() any { + return hos.Conn +} + func (hos *HTTPObfsServer) Read(b []byte) (int, error) { if hos.buf != nil { n := copy(b, hos.buf[hos.offset:]) diff --git a/transport/simple-obfs/tls_server.go b/transport/simple-obfs/tls_server.go index 641002da..a9be2923 100644 --- a/transport/simple-obfs/tls_server.go +++ b/transport/simple-obfs/tls_server.go @@ -19,6 +19,10 @@ type TLSObfsServer struct { firstResponse bool } +func (tos *TLSObfsServer) Upstream() any { + return tos.Conn +} + func (tos *TLSObfsServer) read(b []byte, discardN int) (int, error) { buf := pool.Get(discardN) _, err := io.ReadFull(tos.Conn, buf)