mirror of
https://github.com/MetaCubeX/mihomo.git
synced 2026-10-10 04:03:11 +08:00
chore: better UserFromConn implementation
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:])
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user