Files
Xray-core/transport/internet/masque/hub_test.go
T

151 lines
4.4 KiB
Go

package masque
import (
"context"
"io"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls"
)
func TestServerVersions(t *testing.T) {
for _, c := range []struct {
alpn []string
h2, h3 bool
}{
{alpn: nil, h3: true},
{alpn: []string{"h3"}, h3: true},
{alpn: []string{"h2"}, h2: true},
{alpn: []string{"h2", "http/1.1"}, h2: true},
{alpn: []string{"h3", "h2"}, h2: true, h3: true},
{alpn: []string{"http/1.1"}, h3: true},
} {
h2, h3 := serverVersions(&tls.Config{NextProtocol: c.alpn})
if h2 != c.h2 || h3 != c.h3 {
t.Errorf("serverVersions(%q) = %v, %v, want %v, %v", c.alpn, h2, h3, c.h2, c.h3)
}
}
}
func TestPathMatcher(t *testing.T) {
for _, c := range []struct {
path string
request string
want bool
}{
{DefaultPath, "/.well-known/masque/ip/*/*/", true},
{DefaultPath, "/.well-known/masque/ip/%2A/%2A/", true},
{DefaultPath, "/.well-known/masque/ip/*/*", false},
{DefaultPath, "/.well-known/masque/ip/*/*/?x=1", false},
{DefaultPath, "/.well-known/masque/ip/192.0.2.1/6/", false},
{"/masque?target=*&ipproto=*", "/masque?ipproto=*&target=*", true},
{"/masque?target=*&ipproto=*", "/masque?target=*", false},
} {
m, err := newPathMatcher(c.path)
require.NoError(t, err)
u, err := url.ParseRequestURI(c.request)
require.NoError(t, err)
if got := m.match(u); got != c.want {
t.Errorf("path %q matching %q = %v, want %v", c.path, c.request, got, c.want)
}
}
_, err := newPathMatcher("masque")
require.Error(t, err)
}
func connectIPRequest(target string) *http.Request {
r := httptest.NewRequest(http.MethodGet, "https://proxy.example"+target, nil)
r.Method = http.MethodConnect
r.Proto, r.ProtoMajor, r.ProtoMinor = "HTTP/2.0", 2, 0
r.Header.Set(":protocol", "connect-ip")
r.Header.Set("Capsule-Protocol", "?1")
r.Body = io.NopCloser(&blockingReader{})
return r
}
type blockingReader struct{}
func (*blockingReader) Read([]byte) (int, error) {
time.Sleep(time.Hour)
return 0, io.EOF
}
func serve(t *testing.T, r *http.Request, handle func(*ServerConn)) *httptest.ResponseRecorder {
t.Helper()
path, err := newPathMatcher(DefaultPath)
require.NoError(t, err)
l := &Listener{path: path, addConn: func(conn stat.Connection) {
go func() {
handle(conn.(*ServerConn))
conn.Close()
}()
}}
l.ctx, l.cancel = context.WithCancel(context.Background())
defer l.cancel()
w := httptest.NewRecorder()
done := make(chan struct{})
go func() {
l.ServeHTTP(w, r)
close(done)
}()
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("ServeHTTP did not return")
}
return w
}
func TestListenerRejectsOtherRequests(t *testing.T) {
unexpected := func(*ServerConn) { t.Error("an invalid request reached the proxy") }
require.Equal(t, http.StatusNotFound, serve(t, connectIPRequest("/other"), unexpected).Code)
get := connectIPRequest(DefaultPath)
get.Method = http.MethodGet
require.Equal(t, http.StatusMethodNotAllowed, serve(t, get, unexpected).Code)
websocket := connectIPRequest(DefaultPath)
websocket.Header.Set(":protocol", "websocket")
require.Equal(t, http.StatusNotImplemented, serve(t, websocket, unexpected).Code)
noCapsules := connectIPRequest(DefaultPath)
noCapsules.Header.Del("Capsule-Protocol")
require.Equal(t, http.StatusBadRequest, serve(t, noCapsules, unexpected).Code)
}
func TestServerConnAnswers(t *testing.T) {
t.Run("reject", func(t *testing.T) {
w := serve(t, connectIPRequest(DefaultPath), func(conn *ServerConn) {
require.Equal(t, "connect-ip", conn.Request().Header.Get(":protocol"))
conn.Reject(http.StatusUnauthorized, http.Header{"WWW-Authenticate": {"Basic"}})
})
require.Equal(t, http.StatusUnauthorized, w.Code)
require.Equal(t, "Basic", w.Header().Get("WWW-Authenticate"))
})
t.Run("no answer", func(t *testing.T) {
w := serve(t, connectIPRequest(DefaultPath), func(*ServerConn) {})
require.Equal(t, http.StatusInternalServerError, w.Code)
})
t.Run("accept", func(t *testing.T) {
w := serve(t, connectIPRequest(DefaultPath), func(conn *ServerConn) {
ipConn, err := conn.Accept()
require.NoError(t, err)
require.NotNil(t, ipConn)
_, err = conn.Accept()
require.Error(t, err)
conn.Reject(http.StatusForbidden, nil)
})
require.Equal(t, http.StatusOK, w.Code)
require.Equal(t, "?1", w.Header().Get("Capsule-Protocol"))
})
}