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

681 lines
16 KiB
Go

package masque
import (
"bufio"
"bytes"
"context"
go_errors "errors"
"io"
"maps"
"math"
"net"
"net/http"
"net/url"
"slices"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
"golang.org/x/net/http2"
"golang.org/x/net/http2/hpack"
)
const (
http2MaxConcurrentStreams = 100
http2HandshakeTimeout = 10 * time.Second
http2DefaultHeaderTable = 4096
)
var (
errHTTP2BadPreface = go_errors.New("http2: invalid connection preface")
errHTTP2BadRequest = go_errors.New("http2: malformed request")
errHTTP2StreamClosed = go_errors.New("http2: stream closed")
errHTTP2RequestBodyClosed = go_errors.New("http2: request body closed")
)
type connAddrsKey struct{}
type connAddrs struct {
local net.Addr
remote net.Addr
}
type http2ServerConn struct {
conn net.Conn
handler http.Handler
ctx context.Context
cancel context.CancelFunc
wmu sync.Mutex
bw *bufio.Writer
fr *http2.Framer
hbuf bytes.Buffer
henc *hpack.Encoder
lastFrame atomic.Int64
mu sync.Mutex
cond sync.Cond
err error
maxFrameSize uint32
initialWindow int64
connSendWindow int64
connRecvWindow int64
connRecvUnacked int64
streams map[uint32]*http2ServerStream
lastStreamID uint32
}
type http2ServerStream struct {
c *http2ServerConn
id uint32
ctx context.Context
cancel context.CancelFunc
header http.Header
out *connectip.StreamBuffer
sent chan struct{}
sendWindow int64
recvWindow int64
recvUnacked int64
recv bytes.Buffer
recvEnd bool
wroteHeader bool
sentEnd bool
resetErr error
}
func serveHTTP2(ctx context.Context, conn net.Conn, handler http.Handler) {
ctx, cancel := context.WithCancel(ctx)
c := &http2ServerConn{
conn: conn,
handler: handler,
ctx: context.WithValue(ctx, connAddrsKey{}, connAddrs{local: conn.LocalAddr(), remote: conn.RemoteAddr()}),
cancel: cancel,
bw: bufio.NewWriter(conn),
maxFrameSize: http2DefaultFrameSize,
initialWindow: http2DefaultWindow,
connSendWindow: http2DefaultWindow,
connRecvWindow: http2ConnectionWindow,
streams: make(map[uint32]*http2ServerStream),
}
c.cond.L = &c.mu
c.henc = hpack.NewEncoder(&c.hbuf)
c.henc.SetMaxDynamicTableSizeLimit(0)
c.lastFrame.Store(time.Now().UnixNano())
br := bufio.NewReader(conn)
preface := make([]byte, len(http2.ClientPreface))
conn.SetReadDeadline(time.Now().Add(http2HandshakeTimeout))
if _, err := io.ReadFull(br, preface); err != nil || string(preface) != http2.ClientPreface {
c.fail(errHTTP2BadPreface)
return
}
conn.SetReadDeadline(time.Time{})
c.fr = http2.NewFramer(c.bw, br)
c.fr.SetMaxReadFrameSize(http2DefaultFrameSize)
c.fr.ReadMetaHeaders = hpack.NewDecoder(http2DefaultHeaderTable, nil)
c.fr.MaxHeaderListSize = http2MaxHeaderListSize
if err := c.write(func(fr *http2.Framer) error {
if err := fr.WriteSettings(
http2.Setting{ID: http2.SettingMaxConcurrentStreams, Val: http2MaxConcurrentStreams},
http2.Setting{ID: http2.SettingInitialWindowSize, Val: http2StreamWindow},
http2.Setting{ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize},
http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1},
); err != nil {
return err
}
return fr.WriteWindowUpdate(0, http2ConnectionWindow-http2DefaultWindow)
}); err != nil {
c.fail(err)
return
}
go c.keepAlive()
c.readLoop()
}
func (c *http2ServerConn) write(f func(*http2.Framer) error) error {
c.wmu.Lock()
defer c.wmu.Unlock()
if err := f(c.fr); err != nil {
return err
}
return c.bw.Flush()
}
func (c *http2ServerConn) fail(err error) {
c.mu.Lock()
if c.err == nil {
c.err = err
}
streams := slices.Collect(maps.Values(c.streams))
c.cond.Broadcast()
c.mu.Unlock()
for _, st := range streams {
st.out.CloseWithError(err)
st.cancel()
}
c.cancel()
c.conn.Close()
}
func (c *http2ServerConn) keepAlive() {
ticker := time.NewTicker(http2KeepAlivePeriod)
defer ticker.Stop()
for {
select {
case <-c.ctx.Done():
return
case <-ticker.C:
}
idle := time.Since(time.Unix(0, c.lastFrame.Load()))
if idle >= http2IdleTimeout {
c.fail(errHTTP2IdleTimeout)
return
}
if idle >= http2KeepAlivePeriod {
go c.write(func(fr *http2.Framer) error {
return fr.WritePing(false, [8]byte{})
})
}
}
}
func (c *http2ServerConn) readLoop() {
for {
f, err := c.fr.ReadFrame()
if err != nil {
var streamErr http2.StreamError
if go_errors.As(err, &streamErr) {
if err := c.resetStream(streamErr); err != nil {
c.fail(err)
return
}
continue
}
c.fail(err)
return
}
c.lastFrame.Store(time.Now().UnixNano())
if err := c.handleFrame(f); err != nil {
c.fail(err)
return
}
}
}
func (c *http2ServerConn) handleFrame(f http2.Frame) error {
switch f := f.(type) {
case *http2.SettingsFrame:
if f.IsAck() {
return nil
}
if err := c.applySettings(f); err != nil {
return err
}
return c.write((*http2.Framer).WriteSettingsAck)
case *http2.PingFrame:
if f.IsAck() {
return nil
}
return c.write(func(fr *http2.Framer) error {
return fr.WritePing(true, f.Data)
})
case *http2.WindowUpdateFrame:
return c.handleWindowUpdate(f)
case *http2.MetaHeadersFrame:
return c.handleHeaders(f)
case *http2.DataFrame:
return c.handleData(f)
case *http2.RSTStreamFrame:
c.abortStream(f.StreamID, http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode}, false)
case *http2.PushPromiseFrame:
return http2.ConnectionError(http2.ErrCodeProtocol)
}
return nil
}
func (c *http2ServerConn) applySettings(f *http2.SettingsFrame) error {
c.mu.Lock()
defer c.mu.Unlock()
defer c.cond.Broadcast()
return f.ForeachSetting(func(s http2.Setting) error {
if err := s.Valid(); err != nil {
return err
}
switch s.ID {
case http2.SettingMaxFrameSize:
c.maxFrameSize = s.Val
case http2.SettingInitialWindowSize:
delta := int64(s.Val) - c.initialWindow
c.initialWindow = int64(s.Val)
for _, st := range c.streams {
st.sendWindow += delta
if st.sendWindow > math.MaxInt32 {
return http2.ConnectionError(http2.ErrCodeFlowControl)
}
}
}
return nil
})
}
func (c *http2ServerConn) handleWindowUpdate(f *http2.WindowUpdateFrame) error {
c.mu.Lock()
defer c.mu.Unlock()
if f.StreamID == 0 {
c.connSendWindow += int64(f.Increment)
if c.connSendWindow > math.MaxInt32 {
return http2.ConnectionError(http2.ErrCodeFlowControl)
}
} else if st := c.streams[f.StreamID]; st != nil {
st.sendWindow += int64(f.Increment)
if st.sendWindow > math.MaxInt32 {
return http2.ConnectionError(http2.ErrCodeFlowControl)
}
}
c.cond.Broadcast()
return nil
}
func (c *http2ServerConn) handleHeaders(f *http2.MetaHeadersFrame) error {
c.mu.Lock()
if st := c.streams[f.StreamID]; st != nil {
if !f.StreamEnded() {
c.mu.Unlock()
return http2.ConnectionError(http2.ErrCodeProtocol)
}
st.recvEnd = true
c.cond.Broadcast()
c.mu.Unlock()
return nil
}
if f.StreamID%2 == 0 || f.StreamID <= c.lastStreamID {
c.mu.Unlock()
return http2.ConnectionError(http2.ErrCodeProtocol)
}
c.lastStreamID = f.StreamID
refused := len(c.streams) >= http2MaxConcurrentStreams
c.mu.Unlock()
if refused {
return c.write(func(fr *http2.Framer) error {
return fr.WriteRSTStream(f.StreamID, http2.ErrCodeRefusedStream)
})
}
req, err := newHTTP2Request(f)
if err != nil {
return c.write(func(fr *http2.Framer) error {
return fr.WriteRSTStream(f.StreamID, http2.ErrCodeProtocol)
})
}
ctx, cancel := context.WithCancel(c.ctx)
st := &http2ServerStream{
c: c,
id: f.StreamID,
ctx: ctx,
cancel: cancel,
header: make(http.Header),
out: connectip.NewStreamBuffer(),
sent: make(chan struct{}),
recvWindow: http2StreamWindow,
recvEnd: f.StreamEnded(),
}
req = req.WithContext(ctx)
req.RemoteAddr = c.conn.RemoteAddr().String()
if st.recvEnd {
req.Body = http.NoBody
req.ContentLength = 0
} else {
req.Body = &http2RequestBody{st: st}
req.ContentLength = -1
}
c.mu.Lock()
if c.err != nil {
err := c.err
c.mu.Unlock()
cancel()
return err
}
st.sendWindow = c.initialWindow
c.streams[f.StreamID] = st
c.mu.Unlock()
go c.serveStream(st, req)
return nil
}
func newHTTP2Request(f *http2.MetaHeadersFrame) (*http.Request, error) {
method := f.PseudoValue("method")
scheme := f.PseudoValue("scheme")
authority := f.PseudoValue("authority")
path := f.PseudoValue("path")
protocol := f.PseudoValue("protocol")
if method == "" || (protocol != "" && method != http.MethodConnect) {
return nil, errHTTP2BadRequest
}
var u *url.URL
if method == http.MethodConnect && protocol == "" {
if authority == "" || path != "" || scheme != "" {
return nil, errHTTP2BadRequest
}
u = &url.URL{Host: authority}
} else {
if path == "" || scheme == "" {
return nil, errHTTP2BadRequest
}
var err error
if u, err = url.ParseRequestURI(path); err != nil {
return nil, errHTTP2BadRequest
}
}
header := make(http.Header)
for _, hf := range f.RegularFields() {
header.Add(hf.Name, hf.Value)
}
if protocol != "" {
header.Set(":protocol", protocol)
}
if authority == "" {
authority = header.Get("Host")
}
return &http.Request{
Method: method,
URL: u,
Proto: "HTTP/2.0",
ProtoMajor: 2,
Header: header,
Host: authority,
RequestURI: path,
}, nil
}
func (c *http2ServerConn) handleData(f *http2.DataFrame) error {
size := int64(f.Length)
c.mu.Lock()
c.connRecvWindow -= size
if c.connRecvWindow < 0 {
c.mu.Unlock()
return http2.ConnectionError(http2.ErrCodeFlowControl)
}
st := c.streams[f.StreamID]
if st == nil || st.resetErr != nil || st.recvEnd {
idle := st == nil && f.StreamID > c.lastStreamID
c.connRecvWindow += size
c.mu.Unlock()
if idle {
return http2.ConnectionError(http2.ErrCodeProtocol)
}
if size == 0 {
return nil
}
return c.write(func(fr *http2.Framer) error {
return fr.WriteWindowUpdate(0, uint32(size))
})
}
st.recvWindow -= size
if st.recvWindow < 0 {
c.mu.Unlock()
return http2.ConnectionError(http2.ErrCodeFlowControl)
}
st.recv.Write(f.Data())
padding := size - int64(len(f.Data()))
st.recvUnacked += padding
c.connRecvUnacked += padding
if f.StreamEnded() {
st.recvEnd = true
}
c.cond.Broadcast()
c.mu.Unlock()
return nil
}
func (c *http2ServerConn) resetStream(streamErr http2.StreamError) error {
c.mu.Lock()
if streamErr.StreamID%2 == 1 && streamErr.StreamID > c.lastStreamID {
c.lastStreamID = streamErr.StreamID
}
c.mu.Unlock()
c.abortStream(streamErr.StreamID, streamErr, false)
return c.write(func(fr *http2.Framer) error {
return fr.WriteRSTStream(streamErr.StreamID, streamErr.Code)
})
}
func (c *http2ServerConn) abortStream(id uint32, err error, reset bool) {
c.mu.Lock()
st := c.streams[id]
if st == nil || st.resetErr != nil {
c.mu.Unlock()
return
}
st.resetErr = err
c.cond.Broadcast()
c.mu.Unlock()
st.out.CloseWithError(err)
st.cancel()
if reset {
go c.write(func(fr *http2.Framer) error {
return fr.WriteRSTStream(id, http2.ErrCodeCancel)
})
}
}
func (c *http2ServerConn) serveStream(st *http2ServerStream, req *http.Request) {
go st.sendLoop()
c.handler.ServeHTTP(&http2ResponseWriter{st: st}, req)
st.writeHeader(http.StatusOK)
st.out.Close()
<-st.sent
c.mu.Lock()
finish := st.resetErr == nil && !st.sentEnd && c.err == nil
refuse := finish && !st.recvEnd
st.sentEnd = true
delete(c.streams, st.id)
c.cond.Broadcast()
c.mu.Unlock()
st.cancel()
if finish {
if err := c.write(func(fr *http2.Framer) error {
if err := fr.WriteData(st.id, true, nil); err != nil {
return err
}
if refuse {
return fr.WriteRSTStream(st.id, http2.ErrCodeNo)
}
return nil
}); err != nil {
c.fail(err)
}
}
}
func (st *http2ServerStream) writeHeader(code int) {
c := st.c
c.mu.Lock()
if st.wroteHeader || st.resetErr != nil || c.err != nil {
c.mu.Unlock()
return
}
st.wroteHeader = true
header := st.header.Clone()
maxFrameSize := int(c.maxFrameSize)
c.mu.Unlock()
if err := c.write(func(fr *http2.Framer) error {
c.hbuf.Reset()
c.henc.WriteField(hpack.HeaderField{Name: ":status", Value: strconv.Itoa(code)})
for _, k := range slices.Sorted(maps.Keys(header)) {
name := strings.ToLower(k)
switch name {
case "connection", "proxy-connection", "keep-alive", "transfer-encoding", "upgrade":
continue
}
for _, v := range header[k] {
c.henc.WriteField(hpack.HeaderField{Name: name, Value: v})
}
}
block := c.hbuf.Bytes()
for first := true; first || len(block) > 0; first = false {
chunk := block[:min(len(block), maxFrameSize)]
block = block[len(chunk):]
var err error
if first {
err = fr.WriteHeaders(http2.HeadersFrameParam{StreamID: st.id, BlockFragment: chunk, EndHeaders: len(block) == 0})
} else {
err = fr.WriteContinuation(st.id, len(block) == 0, chunk)
}
if err != nil {
return err
}
}
return nil
}); err != nil {
c.fail(err)
}
}
func (st *http2ServerStream) sendLoop() {
defer close(st.sent)
c := st.c
buf := make([]byte, http2DefaultFrameSize)
for {
n, err := st.out.Read(buf)
for data := buf[:n]; len(data) > 0; {
allowed, err := st.awaitSendWindow(len(data))
if err != nil {
st.out.CloseWithError(err)
return
}
if err := c.write(func(fr *http2.Framer) error {
return fr.WriteData(st.id, false, data[:allowed])
}); err != nil {
c.fail(err)
return
}
data = data[allowed:]
}
if err != nil {
return
}
}
}
func (st *http2ServerStream) awaitSendWindow(n int) (int, error) {
c := st.c
c.mu.Lock()
defer c.mu.Unlock()
for {
switch {
case st.resetErr != nil:
return 0, st.resetErr
case c.err != nil:
return 0, c.err
case st.sentEnd:
return 0, errHTTP2StreamClosed
}
if window := min(c.connSendWindow, st.sendWindow); window > 0 {
n = int(min(int64(n), window, int64(c.maxFrameSize)))
c.connSendWindow -= int64(n)
st.sendWindow -= int64(n)
return n, nil
}
c.cond.Wait()
}
}
type http2ResponseWriter struct {
st *http2ServerStream
}
func (w *http2ResponseWriter) Header() http.Header { return w.st.header }
func (w *http2ResponseWriter) WriteHeader(code int) { w.st.writeHeader(code) }
func (w *http2ResponseWriter) Write(p []byte) (int, error) {
w.st.writeHeader(http.StatusOK)
return w.st.out.Write(p)
}
func (w *http2ResponseWriter) Flush() { w.st.writeHeader(http.StatusOK) }
func (w *http2ResponseWriter) SetWriteDeadline(t time.Time) error {
return w.st.out.SetWriteDeadline(t)
}
type http2RequestBody struct {
st *http2ServerStream
}
func (b *http2RequestBody) Read(p []byte) (int, error) {
st := b.st
c := st.c
c.mu.Lock()
for st.recv.Len() == 0 && !st.recvEnd && st.resetErr == nil && c.err == nil {
c.cond.Wait()
}
if st.recv.Len() == 0 {
err := io.EOF
switch {
case st.resetErr != nil:
err = st.resetErr
case c.err != nil && !st.recvEnd:
err = c.err
}
c.mu.Unlock()
return 0, err
}
n, _ := st.recv.Read(p)
st.recvUnacked += int64(n)
c.connRecvUnacked += int64(n)
var streamUpdate, connUpdate int64
if st.recvUnacked >= http2WindowUpdateSize && !st.recvEnd {
streamUpdate = st.recvUnacked
st.recvUnacked = 0
st.recvWindow += streamUpdate
}
if c.connRecvUnacked >= http2WindowUpdateSize {
connUpdate = c.connRecvUnacked
c.connRecvUnacked = 0
c.connRecvWindow += connUpdate
}
c.mu.Unlock()
if streamUpdate > 0 || connUpdate > 0 {
if err := c.write(func(fr *http2.Framer) error {
if connUpdate > 0 {
if err := fr.WriteWindowUpdate(0, uint32(connUpdate)); err != nil {
return err
}
}
if streamUpdate > 0 {
return fr.WriteWindowUpdate(st.id, uint32(streamUpdate))
}
return nil
}); err != nil {
c.fail(err)
}
}
return n, nil
}
func (b *http2RequestBody) Close() error {
st := b.st
c := st.c
c.mu.Lock()
done := st.recvEnd
c.mu.Unlock()
if !done {
c.abortStream(st.id, errHTTP2RequestBodyClosed, true)
}
return nil
}