diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 393bca9a..a3e648f3 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -599,6 +599,34 @@ jobs: echo "expected all four tests and all 20 arms (8 of them io_uring) to PASS with no SKIP line" exit 1 fi + # celeris#733: the request strings a hijacking handler keeps must not + # read another connection's bytes. The root race step above runs + # without -v, where the io_uring arm of TestHijackKeepsRequestViews + # drops out, and the multishot test skips, without a word when the + # probe finds no ring or no multishot receive. So both run here by + # name with CELERIS_REQUIRE_IOURING_WORKERS=1, and an exact tally: the + # two top-level tests and the four arms (std, epoll, epoll-async, + # io_uring) must PASS, with no SKIP line, and go test must exit 0. + - name: celeris#733 hijacked request strings (io_uring arms required) + if: ${{ !cancelled() }} + shell: bash + env: + CELERIS_REQUIRE_IOURING_WORKERS: "1" + run: | + set -o pipefail + echo "memlock (KiB): $(ulimit -l)" + failed=0 + tests='TestHijackKeepsRequestViews|TestHijackCopiesRequestValuesUnderMultishotRecv' + go test -race -count=1 -timeout=300s -v -run "^(${tests})\$" . 2>&1 | tee /tmp/c733.log || failed=1 + top=$(grep -cE "^--- PASS: (${tests}) \(" /tmp/c733.log || true) + arms=$(grep -cE "^ --- PASS: TestHijackKeepsRequestViews/" /tmp/c733.log || true) + uring=$(grep -cE "^ --- PASS: TestHijackKeepsRequestViews/io_uring" /tmp/c733.log || true) + skipped=$(grep -cE '^[[:space:]]*--- SKIP' /tmp/c733.log || true) + echo "celeris#733 tests: top-level PASS $top (want 2), arm PASS $arms (want 4), io_uring arm PASS $uring (want 1), SKIP lines $skipped, go test failed $failed (want 0)" + if [ "$top" -ne 2 ] || [ "$arms" -ne 4 ] || [ "$uring" -ne 1 ] || [ "$skipped" -ne 0 ] || [ "$failed" -ne 0 ]; then + echo "expected both tests and all four arms (one io_uring) to PASS with no SKIP line" + exit 1 + fi adaptive: name: Adaptive engine (./adaptive/...) diff --git a/context_response.go b/context_response.go index 36f136eb..f77da853 100644 --- a/context_response.go +++ b/context_response.go @@ -1256,6 +1256,17 @@ func (c *Context) BytesWritten() int { // owns the connection and is responsible for closing it. Supported on all // engines for HTTP/1.1 connections. HTTP/2 connections cannot be hijacked // (multiplexed streams share a single TCP connection). +// +// The request's strings stay valid after Hijack, for the handler and for +// any goroutine it hands them to (celeris#733). On epoll and io_uring they +// are views of the connection's receive buffer, which the engine does not +// give to another connection once the connection is hijacked. Hijack also +// copies the request values the Context holds, as [Context.Detach] does, so +// the path, params, headers, query, cookies and body read from the Context +// after Hijack are copies. In io_uring's opt-in multishot receive mode +// (CELERIS_IOURING_MULTISHOT_RECV=1) the request is received into a buffer +// the engine hands back to the kernel when the handler returns: there, keep +// only strings read after Hijack, or clone the ones read before it. func (c *Context) Hijack() (net.Conn, error) { if c.written { return nil, errors.New("celeris: cannot hijack after response written") @@ -1264,6 +1275,9 @@ func (c *Context) Hijack() (net.Conn, error) { if !ok { return nil, ErrHijackNotSupported } + // Before the engine lets go of the connection, and so of its receive + // buffer (celeris#733). + c.cloneRequestValues() conn, err := h.Hijack(c.stream) if err != nil { return nil, err @@ -1420,25 +1434,7 @@ func (c *Context) Detach() (done func()) { if c.detached { return func() {} // already detached — return no-op done } - // Materialize any unsafe string headers (zero-copy H1 headers backed by - // the connection's read buffer) before the handler returns and the buffer - // is reused for the next recv. Pseudo-header keys are string literals - // (safe), but their values (:authority, :path for non-"/" paths, :method - // for non-standard methods) may be UnsafeString backed by the buffer. - c.stream.MaterializeHeaders() - for i, h := range c.stream.Headers { - if len(h[0]) > 0 && h[0][0] == ':' { - c.stream.Headers[i][1] = strings.Clone(h[1]) - continue - } - c.stream.Headers[i][0] = strings.Clone(h[0]) - c.stream.Headers[i][1] = strings.Clone(h[1]) - } - // Also materialize extracted fields that may reference the buffer. - c.method = strings.Clone(c.method) - c.path = strings.Clone(c.path) - c.rawQuery = strings.Clone(c.rawQuery) - c.materializeRequestViews() + c.cloneRequestValues() c.extended = true c.detached = true @@ -1462,9 +1458,36 @@ func (c *Context) Detach() (done func()) { } } -// materializeRequestViews clones, for [Context.Detach], the request-derived -// strings the Context keeps outside the header slice and the method, path and -// raw query: the route params, the parsed query and cookie caches, the H1 +// cloneRequestValues copies every request value the Context holds that can be +// a view of the engine's receive buffer, for [Context.Detach] and +// [Context.Hijack]: after either, the handler, or a goroutine it starts, can +// keep using those values after the engine has received into, or given +// away, the buffer they view. +func (c *Context) cloneRequestValues() { + // Materialize any unsafe string headers (zero-copy H1 headers backed by + // the connection's read buffer) before the handler returns and the buffer + // is reused for the next recv. Pseudo-header keys are string literals + // (safe), but their values (:authority, :path for non-"/" paths, :method + // for non-standard methods) may be UnsafeString backed by the buffer. + c.stream.MaterializeHeaders() + for i, h := range c.stream.Headers { + if len(h[0]) > 0 && h[0][0] == ':' { + c.stream.Headers[i][1] = strings.Clone(h[1]) + continue + } + c.stream.Headers[i][0] = strings.Clone(h[0]) + c.stream.Headers[i][1] = strings.Clone(h[1]) + } + // Also materialize extracted fields that may reference the buffer. + c.method = strings.Clone(c.method) + c.path = strings.Clone(c.path) + c.rawQuery = strings.Clone(c.rawQuery) + c.materializeRequestViews() +} + +// materializeRequestViews clones, for [Context.Detach] and [Context.Hijack], +// the request-derived strings the Context keeps outside the header slice and +// the method, path and raw query: the route params, the parsed query and cookie caches, the H1 // Host, the strings middleware store on the Context from request headers // (request ID, client-IP/host/scheme overrides, SetString values), the // request body, and the response headers set so far. diff --git a/engine/epoll/loop.go b/engine/epoll/loop.go index c9af519e..00041c5b 100644 --- a/engine/epoll/loop.go +++ b/engine/epoll/loop.go @@ -1779,6 +1779,15 @@ func (l *Loop) hijackConn(fd int) (net.Conn, error) { // drainDetachQueue, which also runs there. l.removeDirty(cs) l.dropAsk(cs) // celeris#657 P8: never pool a connState an ask still names + // The request's strings are views of cs.buf (celeris#733). The + // handler is still running, and a hijacking handler typically keeps + // the path, params and headers for the goroutine that will serve + // the connection, so the buffer must not go back to the pool with + // cs: the next connection to take cs from the pool would receive its + // request into it, on any worker, even before this handler returns. + // Dropping it costs one buffer allocation per hijack, when that + // connState is next acquired. + cs.buf = nil releaseConnState(cs) } return c, err diff --git a/engine/iouring/worker.go b/engine/iouring/worker.go index b734d332..71933c01 100644 --- a/engine/iouring/worker.go +++ b/engine/iouring/worker.go @@ -2245,9 +2245,20 @@ func (w *Worker) hijackConn(fd int) (net.Conn, error) { // STEAL the first bytes the hijacker tries to read. Cancel it by its // generation-tagged user_data and defer the pool release until the // terminal CQE arrives, exactly like finishClose. + // + // Then drop cs instead of recycling it (celeris#733): the request's + // strings are views of cs.buf (or of cs.detectAccum, when the request + // began in an earlier recv), and a hijacking handler typically keeps + // the path, params and headers for the goroutine that serves the + // connection. Recycled, cs would hand those buffers to the next + // connection this worker accepts, which would receive its request into + // them. The detached release holds cs until the kernel is done with + // cs.buf and then leaves it to the garbage collector, which frees the + // buffers once the last view of them is gone. The cost is one connState + // allocation per hijack. w.cancelConnOps(fd, cs) w.noteClosedInflight(cs) - w.queuePendingRelease(cs) + w.queuePendingReleaseDetached(cs) f := os.NewFile(uintptr(fd), "tcp") c, err := net.FileConn(f) _ = f.Close() diff --git a/hijack_keeps_request_views_linux_test.go b/hijack_keeps_request_views_linux_test.go new file mode 100644 index 00000000..2954a9fe --- /dev/null +++ b/hijack_keeps_request_views_linux_test.go @@ -0,0 +1,287 @@ +//go:build linux + +package celeris_test + +import ( + "bufio" + "fmt" + "io" + "net" + "net/http" + "os" + "strconv" + "strings" + "testing" + "time" + + "github.com/goceleris/celeris" +) + +// c733Get sends one GET for path on conn, with an Authorization header, and +// returns the response body. +func c733Get(conn net.Conn, br *bufio.Reader, path, auth string) (string, error) { + if _, err := conn.Write([]byte("GET " + path + " HTTP/1.1\r\nHost: x\r\nAuthorization: " + auth + "\r\n\r\n")); err != nil { + return "", err + } + resp, err := http.ReadResponse(br, nil) + if err != nil { + return "", err + } + body, err := io.ReadAll(resp.Body) + _ = resp.Body.Close() + if err != nil { + return "", err + } + if resp.StatusCode != http.StatusOK { + return "", fmt.Errorf("status %d, body %q", resp.StatusCode, body) + } + return string(body), nil +} + +// c733DialWorker dials addr until the connection is served by the worker +// whose ID is want (GET /w answers with c.WorkerID()), and returns it. A +// worker ID of -1 (std) accepts the first connection. The connections that +// land on another worker are closed. +func c733DialWorker(t *testing.T, addr string, want int, auth string) (net.Conn, *bufio.Reader) { + t.Helper() + for try := 0; try < 64; try++ { + conn, err := net.DialTimeout("tcp", addr, 2*time.Second) + if err != nil { + t.Fatal(err) + } + _ = conn.SetDeadline(time.Now().Add(60 * time.Second)) + br := bufio.NewReader(conn) + body, err := c733Get(conn, br, "/w", auth) + if err != nil { + t.Fatalf("GET /w: %v", err) + } + if got, _ := strconv.Atoi(body); got == want || want < 0 { + return conn, br + } + _ = conn.Close() + } + t.Fatalf("64 connections, none served by worker %d", want) + return nil, nil +} + +const c733Secret = "SECRETSECRETSECRETSECRETSECRETSECRET" + +// TestHijackKeepsRequestViews pins celeris#733. +// +// On epoll and io_uring the request strings a handler reads from the Context +// (the path, the route params, the headers) are views of the connection's +// receive buffer. Context.Hijack hands the connection to the handler, and the +// engine gave the connection's state, receive buffer included, back to its +// pool: epoll inside the Hijack call, io_uring once the cancelled receive had +// completed. The next connection the worker accepted took that state and +// received its request into the same buffer, so the strings the hijacking +// handler kept for the goroutine that serves the connection read the other +// connection's bytes: a kept header read another client's Authorization +// value. +// +// Each round, connection A's handler keeps its param, header and path and +// hijacks; then connection B, a new connection that is served by A's worker +// (it is redialled until it is), sends requests whose Authorization value +// lies over the bytes A's strings view. After that, the kept strings must +// still read A's values. std serves the request from copies and is the +// control, as is epoll with the route marked Async: its handler, and so +// Hijack, runs on the connection's dispatch goroutine, and epoll never pools +// the state of a connection hijacked there (celeris#668). io_uring refuses +// Hijack on an async worker (celeris#539), so it has no async arm. +func TestHijackKeepsRequestViews(t *testing.T) { + type arm struct { + name string + engine celeris.EngineType + asyncRoute bool // mark /hj Async: the handler runs on the dispatch goroutine + } + arms := []arm{ + {"std", celeris.Std, false}, + {"epoll", celeris.Epoll, false}, + {"epoll-async", celeris.Epoll, true}, + } + if ok, p := c714ProbeIOUring(); ok { + arms = append(arms, arm{"io_uring", celeris.IOUring, false}) + } else if os.Getenv("CELERIS_REQUIRE_IOURING_WORKERS") == "1" { + t.Fatalf("io_uring tier=%s kernel=%s, and CELERIS_REQUIRE_IOURING_WORKERS=1 forbids dropping the io_uring arm", p.IOUringTier, p.KernelVersion) + } else { + t.Logf("io_uring tier=%s kernel=%s: io_uring arm not run", p.IOUringTier, p.KernelVersion) + } + + type kept struct{ param, header, path string } + type hijacked struct { + k kept + worker int + conn net.Conn + err error + } + for _, a := range arms { + t.Run(a.name, func(t *testing.T) { + got := make(chan hijacked, 1) + addr, stopServer := startC714DetachServer(t, func() *celeris.Server { + srv := celeris.New(celeris.Config{Engine: a.engine, Workers: 2}) + rt := srv.GET("/hj/:id", func(c *celeris.Context) error { + k := kept{param: c.Param("id"), header: c.Header("x-token"), path: c.Path()} + w := c.WorkerID() + conn, err := c.Hijack() + got <- hijacked{k: k, worker: w, conn: conn, err: err} + return nil + }) + if a.asyncRoute { + rt.Async() + } + srv.GET("/w", func(c *celeris.Context) error { return c.String(200, "%d", c.WorkerID()) }) + return srv + }) + defer stopServer() + + const rounds = 20 + wrong, leaked := 0, 0 + var samples []string + for i := 0; i < rounds; i++ { + id := fmt.Sprintf("id%06d", i) + want := kept{param: id, header: "token-" + id, path: "/hj/" + id} + ca, err := net.DialTimeout("tcp", addr, 2*time.Second) + if err != nil { + t.Fatal(err) + } + if _, err := ca.Write([]byte("GET /hj/" + id + " HTTP/1.1\r\nHost: x\r\nX-Token: token-" + id + "\r\n\r\n")); err != nil { + t.Fatal(err) + } + var h hijacked + select { + case h = <-got: + case <-time.After(10 * time.Second): + t.Fatalf("round %d: the handler never hijacked", i) + } + if h.err != nil { + t.Fatalf("round %d: Hijack: %v", i, h.err) + } + + // B: a new connection on A's worker, served while A's handler + // keeps its strings. Its Authorization value covers the + // offsets A's strings view. + b, br := c733DialWorker(t, addr, h.worker, c733Secret) + if _, err := c733Get(b, br, "/w", c733Secret); err != nil { + t.Fatalf("round %d: B: %v", i, err) + } + _ = b.Close() + + // What the hijacked session's goroutine would read now. + if h.k != want { + wrong++ + if strings.Contains(h.k.param+h.k.header+h.k.path, "SECRET") { + leaked++ + } + if len(samples) < 4 { + samples = append(samples, fmt.Sprintf("want %q got %q", want, h.k)) + } + } + _ = h.conn.Close() + _ = ca.Close() + } + t.Logf("C733HIJACK arm=%s rounds=%d kept strings wrong=%d, holding B's Authorization bytes=%d", a.name, rounds, wrong, leaked) + if wrong > 0 { + t.Errorf("strings a hijacking handler kept changed after another connection was served: %d/%d (%d with its Authorization bytes); samples %q", + wrong, rounds, leaked, samples) + } + }) + } +} + +// TestHijackCopiesRequestValuesUnderMultishotRecv pins the Context half of +// celeris#733. In io_uring's opt-in multishot receive mode the request is +// received into a buffer of the worker's provided-buffer ring, which the +// engine hands back to the kernel when the handler returns, hijacked or not. +// The engine cannot keep that buffer, so strings read from the Context +// before Hijack are views the kernel will write into, and only copies made +// at Hijack survive. Context.Hijack copies the request values the Context +// holds, as Context.Detach does, so what the handler reads from the Context +// after Hijack must still read its request once the ring has cycled: B, a +// connection on A's worker, sends more requests than the ring has buffers, +// one receive each, so every buffer, A's included, is written again. +// +// The strings read before Hijack are the rig's witness: they are views of +// A's ring buffer, so they must have changed, or the ring did not cycle and +// the test has not shown anything. +func TestHijackCopiesRequestValuesUnderMultishotRecv(t *testing.T) { + ok, p := c714ProbeIOUring() + if !ok || !p.MultishotRecv { + if os.Getenv("CELERIS_REQUIRE_IOURING_WORKERS") == "1" { + t.Fatalf("io_uring tier=%s kernel=%s multishotRecv=%t, and CELERIS_REQUIRE_IOURING_WORKERS=1 forbids skipping", p.IOUringTier, p.KernelVersion, p.MultishotRecv) + } + t.Skipf("io_uring tier=%s kernel=%s multishotRecv=%t: no multishot receive", p.IOUringTier, p.KernelVersion, p.MultishotRecv) + } + t.Setenv("CELERIS_IOURING_MULTISHOT_RECV", "1") + + type kept struct{ param, header, path string } + type hijacked struct { + pre, post kept + worker int + conn net.Conn + err error + } + got := make(chan hijacked, 1) + addr, stopServer := startC714DetachServer(t, func() *celeris.Server { + srv := celeris.New(celeris.Config{Engine: celeris.IOUring, Workers: 2}) + srv.GET("/hj/:id", func(c *celeris.Context) error { + pre := kept{param: c.Param("id"), header: c.Header("x-token"), path: c.Path()} + w := c.WorkerID() + conn, err := c.Hijack() + post := kept{param: c.Param("id"), header: c.Header("x-token"), path: c.Path()} + got <- hijacked{pre: pre, post: post, worker: w, conn: conn, err: err} + return nil + }) + srv.GET("/w", func(c *celeris.Context) error { return c.String(200, "%d", c.WorkerID()) }) + return srv + }) + defer stopServer() + + const id = "id733733" + want := kept{param: id, header: "token-" + id, path: "/hj/" + id} + ca, err := net.DialTimeout("tcp", addr, 2*time.Second) + if err != nil { + t.Fatal(err) + } + defer func() { _ = ca.Close() }() + if _, err := ca.Write([]byte("GET /hj/" + id + " HTTP/1.1\r\nHost: x\r\nX-Token: token-" + id + "\r\n\r\n")); err != nil { + t.Fatal(err) + } + var h hijacked + select { + case h = <-got: + case <-time.After(10 * time.Second): + t.Fatal("the handler never hijacked") + } + if h.err != nil { + t.Fatalf("Hijack: %v", h.err) + } + defer func() { _ = h.conn.Close() }() + + // B: one keep-alive connection on A's worker, one request per round + // trip, so one receive, and one ring buffer, per request. + // CELERIS_IOURING_PBUF_COUNT is unset, so the worker sizes the ring to + // twice its default conns per worker, raised to bufRingCountMin: 1024 + // buffers (engine/iouring resolveBufRingCount). Twice that many requests + // cycle it twice. + b, br := c733DialWorker(t, addr, h.worker, c733Secret) + defer func() { _ = b.Close() }() + const requests = 2048 + for i := 0; i < requests; i++ { + if _, err := c733Get(b, br, "/w", c733Secret); err != nil { + t.Fatalf("B request %d: %v", i, err) + } + } + + // Copies, taken before the server stops: the ring's memory is unmapped + // then, and the views read before Hijack point into it. + pre := strings.Clone(h.pre.param + "|" + h.pre.header + "|" + h.pre.path) + post := strings.Clone(h.post.param + "|" + h.post.header + "|" + h.post.path) + wantS := want.param + "|" + want.header + "|" + want.path + t.Logf("C733MSHOT worker=%d requests=%d before-Hijack %q after-Hijack %q", h.worker, requests, pre, post) + if pre == wantS { + t.Fatalf("witness: the strings read before Hijack still read A's request after %d requests on B, so the ring did not cycle (is multishot receive on?) and this test shows nothing", requests) + } + if post != wantS { + t.Errorf("strings read from the Context after Hijack changed once the ring cycled: %q, want %q: Hijack did not copy the request values (celeris#733)", post, wantS) + } +}