Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
145 changes: 145 additions & 0 deletions engine/epoll/hijack_inline_async_loop_linux_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,145 @@
//go:build linux

package epoll

import (
"bufio"
"context"
"fmt"
"io"
"net"
"net/http"
"sync/atomic"
"testing"
"time"

"github.com/goceleris/celeris/engine"
"github.com/goceleris/celeris/protocol/h2/stream"
"github.com/goceleris/celeris/resource"
)

// A loop with async dispatch on (Config.AsyncHandlers, or any route marked
// Async) runs a connection's requests inline, in InlineMode, until one of
// them reaches an async route. A handler that runs inline there and calls
// Hijack takes hijackConn's inline branch, which returns the connState to the
// pool inside the Hijack call, h1State included (nil after release). drainRead
// then cleared InlineMode through cs.h1State, before it looked at
// ErrHijacked: a nil dereference on the loop goroutine, which took the whole
// process down on the first such request. On a loop whose released connState
// had already been taken by another worker's accept, the same line wrote
// that other connection's parser state.

// inlineHijackHandler: /hj hijacks and answers on the raw conn, /async is the
// async route (when asyncRoutes is set), anything else answers normally.
type inlineHijackHandler struct {
asyncRoutes bool
hijacked *atomic.Int64
}

func (h inlineHijackHandler) HandleStream(_ context.Context, s *stream.Stream) error {
if s.ResponseWriter == nil {
return nil
}
if s.Path != "/hj" {
return s.ResponseWriter.WriteResponse(s, 200,
[][2]string{{"content-type", "text/plain"}, {"content-length", "2"}}, []byte("ok"))
}
hj, ok := s.ResponseWriter.(stream.Hijacker)
if !ok {
return fmt.Errorf("response writer %T cannot hijack", s.ResponseWriter)
}
c, err := hj.Hijack(s)
if err != nil {
return err
}
h.hijacked.Add(1)
_, _ = io.WriteString(c, "HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nhj")
return c.Close()
}
func (h inlineHijackHandler) RouteAsync(_, path string) bool {
return h.asyncRoutes && path == "/async"
}
func (h inlineHijackHandler) HasAsyncRoutes() bool { return h.asyncRoutes }

func TestInlineHijackOnAsyncLoop(t *testing.T) {
Comment thread
FumingPower3925 marked this conversation as resolved.
for _, tc := range []struct {
name string
asyncRoutes bool
}{
// Config.AsyncHandlers with no route marked: no resolver is wired,
// and every request runs inline in InlineMode.
{"async-loop-no-async-routes", false},
// A server with one Async route: /hj is a sync route, run inline.
{"async-loop-sync-route", true},
} {
t.Run(tc.name, func(t *testing.T) {
var hijacked atomic.Int64
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("pick port: %v", err)
}
addr := ln.Addr().String()
_ = ln.Close()
e, err := New(resource.Config{
Addr: addr,
Protocol: engine.HTTP1,
Resources: resource.Resources{Workers: 2},
AsyncHandlers: true,
}, inlineHijackHandler{asyncRoutes: tc.asyncRoutes, hijacked: &hijacked})
if err != nil {
t.Fatalf("epoll engine: %v", err)
}
ctx, cancel := context.WithCancel(context.Background())
errCh := make(chan error, 1)
go func() { errCh <- e.Listen(ctx) }()
defer func() {
cancel()
select {
case <-errCh:
case <-time.After(5 * time.Second):
}
}()
for dl := time.Now().Add(10 * time.Second); e.Addr() == nil && time.Now().Before(dl); {
time.Sleep(10 * time.Millisecond)
}
if e.Addr() == nil {
t.Fatal("engine did not bind")
}
if !e.loops[0].async {
t.Fatal("precondition: the loop is not in async mode")
}

get := func(path string) (string, error) {
c, err := net.DialTimeout("tcp", addr, 3*time.Second)
if err != nil {
return "", err
}
defer func() { _ = c.Close() }()
_ = c.SetDeadline(time.Now().Add(5 * time.Second))
if _, err := fmt.Fprintf(c, "GET %s HTTP/1.1\r\nHost: x\r\n\r\n", path); err != nil {
return "", err
}
resp, err := http.ReadResponse(bufio.NewReader(c), nil)
if err != nil {
return "", err
}
b, err := io.ReadAll(resp.Body)
_ = resp.Body.Close()
return string(b), err
}

const n = 20
for i := 0; i < n; i++ {
if got, err := get("/hj"); err != nil || got != "hj" {
t.Fatalf("hijack %d: body %q, err %v", i, got, err)
}
if got, err := get("/ok"); err != nil || got != "ok" {
t.Fatalf("request after hijack %d: body %q, err %v", i, got, err)
}
}
if got := hijacked.Load(); got != n {
t.Errorf("hijacked %d connections, want %d", got, n)
}
})
}
}
7 changes: 6 additions & 1 deletion engine/epoll/loop.go
Original file line number Diff line number Diff line change
Expand Up @@ -1394,7 +1394,12 @@ func (l *Loop) drainRead(fd int, now int64) {
cs.h1State.InlineMode = true
}
processErr = conn.ProcessH1(cs.ctx, data, cs.h1State, l.handler, writeFn)
if tryInline {
// A handler that ran inline and hijacked has had cs released to
// the pool inside the Hijack call (hijackConn's inline branch):
// cs.h1State is nil, and cs may already be another accept's
// connection. So cs is not touched again once ProcessH1 reports
// the hijack; the ErrHijacked return below is the only way out.
if tryInline && !errors.Is(processErr, conn.ErrHijacked) {
cs.h1State.InlineMode = false
}
if errors.Is(processErr, conn.ErrAsyncDispatch) {
Expand Down
63 changes: 63 additions & 0 deletions hijack_async_handlers_epoll_linux_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,63 @@
//go:build linux

package celeris_test

import (
"bufio"
"io"
"net"
"net/http"
"testing"
"time"

"github.com/goceleris/celeris"
)

// TestHijackWithAsyncHandlersOnEpoll is the user-facing face of the
// engine/epoll TestInlineHijackOnAsyncLoop: with Config.AsyncHandlers, a
// route that inherits the default starts inline (celeris#356), so its
// handler runs on the event loop, and Hijack there released the connState
// that drainRead then dereferenced. The first hijack crashed the process.
func TestHijackWithAsyncHandlersOnEpoll(t *testing.T) {
addr, stopServer := startC714DetachServer(t, func() *celeris.Server {
srv := celeris.New(celeris.Config{Engine: celeris.Epoll, Workers: 2, AsyncHandlers: true})
srv.GET("/hj", func(c *celeris.Context) error {
conn, err := c.Hijack()
if err != nil {
return err
}
_, _ = io.WriteString(conn, "HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nhj")
return conn.Close()
})
srv.GET("/ok", func(c *celeris.Context) error { return c.String(200, "ok") })
return srv
})
defer stopServer()

get := func(path string) (string, error) {
c, err := net.DialTimeout("tcp", addr, 3*time.Second)
if err != nil {
return "", err
}
defer func() { _ = c.Close() }()
_ = c.SetDeadline(time.Now().Add(5 * time.Second))
if _, err := io.WriteString(c, "GET "+path+" HTTP/1.1\r\nHost: x\r\n\r\n"); err != nil {
return "", err
}
resp, err := http.ReadResponse(bufio.NewReader(c), nil)
if err != nil {
return "", err
}
b, err := io.ReadAll(resp.Body)
_ = resp.Body.Close()
return string(b), err
}
for i := 0; i < 20; i++ {
if got, err := get("/hj"); err != nil || got != "hj" {
t.Fatalf("hijack %d: body %q, err %v", i, got, err)
}
if got, err := get("/ok"); err != nil || got != "ok" {
t.Fatalf("request after hijack %d: body %q, err %v", i, got, err)
}
}
}
Loading