Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
753d24be71 |
@@ -3,6 +3,7 @@ package proxy_test
|
|||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
"io"
|
"io"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -166,6 +167,39 @@ func TestWebSocketBytesAreCountedOnceItCloses(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestWebSocketPassesTheAnswerAfterTheClientStopsSending(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
app := startApp(t, echoOnceTheClientStops)
|
||||||
|
addr, out := startProxy(t, app.URL,
|
||||||
|
map[string]string{trustedProxies: trustLocalhost})
|
||||||
|
s := &sender{t: t, addr: addr, out: out}
|
||||||
|
|
||||||
|
conn, reader := s.openWebSocket()
|
||||||
|
send(t, conn, uploadBody)
|
||||||
|
|
||||||
|
// The client closes its sending side and waits for the answer, which the
|
||||||
|
// app sends only once it has seen the client stop. smallwebwaf passes the
|
||||||
|
// close on to the app through CloseWrite on upgradedConn; without that,
|
||||||
|
// it closes both connections, and the answer is lost.
|
||||||
|
tcp, ok := conn.(*net.TCPConn)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("connection is a %T, want a *net.TCPConn", conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
err := tcp.CloseWrite()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("close the sending side: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
got, err := io.ReadAll(reader)
|
||||||
|
if err != nil || string(got) != uploadBody {
|
||||||
|
t.Errorf("got %q (%v), want %q", got, err, uploadBody)
|
||||||
|
}
|
||||||
|
|
||||||
|
s.closeWebSocket(conn)
|
||||||
|
}
|
||||||
|
|
||||||
func TestBytesCountSaysWhichBytesCount(t *testing.T) {
|
func TestBytesCountSaysWhichBytesCount(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -432,12 +466,52 @@ func answerAfterUpgrade(w http.ResponseWriter, _ *http.Request) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// echoOnceTheClientStops is an app that switches protocols, as for a
|
||||||
|
// WebSocket, reads what the client sends until the client stops sending,
|
||||||
|
// and then sends it all back.
|
||||||
|
func echoOnceTheClientStops(w http.ResponseWriter, _ *http.Request) {
|
||||||
|
conn, buffered, err := http.NewResponseController(w).Hijack()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
defer func() {
|
||||||
|
_ = conn.Close()
|
||||||
|
}()
|
||||||
|
|
||||||
|
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
|
||||||
|
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||||
|
_ = buffered.Flush()
|
||||||
|
|
||||||
|
received, _ := io.ReadAll(buffered)
|
||||||
|
_, _ = buffered.Write(received)
|
||||||
|
_ = buffered.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
// webSocket opens a WebSocket from client to answerAfterUpgrade, sends a
|
// webSocket opens a WebSocket from client to answerAfterUpgrade, sends a
|
||||||
// line of bodyBytes on it, reads the answer, and closes it. It checks the
|
// line of bodyBytes on it, reads the answer, and closes it. It checks the
|
||||||
// answer, and the log line as request does, and returns the log line.
|
// answer, and the log line as request does, and returns the log line.
|
||||||
func (s *sender) webSocket() logLine {
|
func (s *sender) webSocket() logLine {
|
||||||
s.t.Helper()
|
s.t.Helper()
|
||||||
|
|
||||||
|
conn, reader := s.openWebSocket()
|
||||||
|
send(s.t, conn, strings.Repeat("u", bodyBytes-1)+"\n")
|
||||||
|
|
||||||
|
got, err := reader.ReadString('\n')
|
||||||
|
if err != nil || len(got) != answerBytes {
|
||||||
|
s.t.Errorf("got %d bytes (%v), want %d", len(got), err, answerBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.closeWebSocket(conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
// openWebSocket sends a request from client to switch protocols, as for a
|
||||||
|
// WebSocket, and checks that the app switches. It returns the connection,
|
||||||
|
// on which reading fails once waitLimit has passed, and a reader of what
|
||||||
|
// the app sends on it.
|
||||||
|
func (s *sender) openWebSocket() (net.Conn, *bufio.Reader) {
|
||||||
|
s.t.Helper()
|
||||||
|
|
||||||
conn := dial(s.t, s.addr)
|
conn := dial(s.t, s.addr)
|
||||||
send(s.t, conn, "GET /socket HTTP/1.1\r\nHost: "+appHost+"\r\n"+forwardedFor+
|
send(s.t, conn, "GET /socket HTTP/1.1\r\nHost: "+appHost+"\r\n"+forwardedFor+
|
||||||
": "+client+"\r\nConnection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
": "+client+"\r\nConnection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||||
@@ -460,12 +534,13 @@ func (s *sender) webSocket() logLine {
|
|||||||
s.t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
|
s.t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
|
||||||
}
|
}
|
||||||
|
|
||||||
send(s.t, conn, strings.Repeat("u", bodyBytes-1)+"\n")
|
return conn, reader
|
||||||
|
}
|
||||||
|
|
||||||
got, err := reader.ReadString('\n')
|
// closeWebSocket closes conn, a WebSocket openWebSocket opened, checks its
|
||||||
if err != nil || len(got) != answerBytes {
|
// log line as request does, and returns it.
|
||||||
s.t.Errorf("got %d bytes (%v), want %d", len(got), err, answerBytes)
|
func (s *sender) closeWebSocket(conn net.Conn) logLine {
|
||||||
}
|
s.t.Helper()
|
||||||
|
|
||||||
_ = conn.Close()
|
_ = conn.Close()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user