Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
753d24be71 |
@@ -3,6 +3,7 @@ package proxy_test
|
||||
import (
|
||||
"bufio"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"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) {
|
||||
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
|
||||
// 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.
|
||||
func (s *sender) webSocket() logLine {
|
||||
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)
|
||||
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")
|
||||
@@ -460,12 +534,13 @@ func (s *sender) webSocket() logLine {
|
||||
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')
|
||||
if err != nil || len(got) != answerBytes {
|
||||
s.t.Errorf("got %d bytes (%v), want %d", len(got), err, answerBytes)
|
||||
}
|
||||
// closeWebSocket closes conn, a WebSocket openWebSocket opened, checks its
|
||||
// log line as request does, and returns it.
|
||||
func (s *sender) closeWebSocket(conn net.Conn) logLine {
|
||||
s.t.Helper()
|
||||
|
||||
_ = conn.Close()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user