diff --git a/httpx/cookie.go b/httpx/cookie.go new file mode 100644 index 0000000..611af48 --- /dev/null +++ b/httpx/cookie.go @@ -0,0 +1,700 @@ +package httpx + +import ( + "bytes" + "errors" + "io" + "path/filepath" + "strconv" + "strings" + "sync" + "time" +) + +var zeroTime time.Time + +var ( + // cookieExpireDelete may be set on Cookie.Expire for expiring the given cookie. + cookieExpireDelete = time.Date(2009, time.November, 10, 23, 0, 0, 0, time.UTC) + + // cookieExpireUnlimited indicates that the cookie doesn't expire. + cookieExpireUnlimited = zeroTime +) + +// CookieSameSite is an enum for the mode in which the SameSite flag should be set for the given cookie. +// See https://tools.ietf.org/html/draft-ietf-httpbis-cookie-same-site-00 for details. +type CookieSameSite int + +const ( + // CookieSameSiteDisabled removes the SameSite flag. + CookieSameSiteDisabled CookieSameSite = iota + // CookieSameSiteDefaultMode sets the SameSite flag. + CookieSameSiteDefaultMode + // CookieSameSiteLaxMode sets the SameSite flag with the "Lax" parameter. + CookieSameSiteLaxMode + // CookieSameSiteStrictMode sets the SameSite flag with the "Strict" parameter. + CookieSameSiteStrictMode + // CookieSameSiteNoneMode sets the SameSite flag with the "None" parameter. + // See https://tools.ietf.org/html/draft-west-cookie-incrementalism-00 + CookieSameSiteNoneMode +) + +// acquireCookie returns an empty Cookie object from the pool. +// +// The returned object may be returned back to the pool with ReleaseCookie. +// This allows reducing GC load. +func acquireCookie() *Cookie { + return cookiePool.Get().(*Cookie) +} + +// releaseCookie returns the Cookie object acquired with AcquireCookie back +// to the pool. +// +// Do not access released Cookie object, otherwise data races may occur. +func releaseCookie(c *Cookie) { + c.Reset() + cookiePool.Put(c) +} + +var cookiePool = &sync.Pool{ + New: func() any { + return &Cookie{} + }, +} + +// Cookie represents HTTP response cookie. +// +// Do not copy Cookie objects. Create new object and use CopyTo instead. +// +// Cookie instance MUST NOT be used from concurrently running goroutines. +type Cookie struct { + noCopy noCopy + + key []byte + value []byte + expire time.Time + maxAge int + domain []byte + path []byte + + httpOnly bool + secure bool + sameSite CookieSameSite + + bufKV argsKV + buf []byte +} + +// CopyTo copies src cookie to c. +func (c *Cookie) CopyTo(src *Cookie) { + c.Reset() + c.key = append(c.key, src.key...) + c.value = append(c.value, src.value...) + c.expire = src.expire + c.maxAge = src.maxAge + c.domain = append(c.domain, src.domain...) + c.path = append(c.path, src.path...) + c.httpOnly = src.httpOnly + c.secure = src.secure + c.sameSite = src.sameSite +} + +// HTTPOnly returns true if the cookie is http only. +func (c *Cookie) HTTPOnly() bool { + return c.httpOnly +} + +// SetHTTPOnly sets cookie's httpOnly flag to the given value. +func (c *Cookie) SetHTTPOnly(httpOnly bool) { + c.httpOnly = httpOnly +} + +// Secure returns true if the cookie is secure. +func (c *Cookie) Secure() bool { + return c.secure +} + +// SetSecure sets cookie's secure flag to the given value. +func (c *Cookie) SetSecure(secure bool) { + c.secure = secure +} + +// SameSite returns the SameSite mode. +func (c *Cookie) SameSite() CookieSameSite { + return c.sameSite +} + +// SetSameSite sets the cookie's SameSite flag to the given value. +// Set value CookieSameSiteNoneMode will set Secure to true also to avoid browser rejection. +func (c *Cookie) SetSameSite(mode CookieSameSite) { + c.sameSite = mode + if mode == CookieSameSiteNoneMode { + c.SetSecure(true) + } +} + +// Path returns cookie path. +func (c *Cookie) Path() []byte { + return c.path +} + +// SetPath sets cookie path. +func (c *Cookie) SetPath(path string) { + c.buf = append(c.buf[:0], path...) + c.path = normalizePath(c.path, b2s(c.buf)) +} + +// SetPathBytes sets cookie path. +func (c *Cookie) SetPathBytes(path []byte) { + c.buf = append(c.buf[:0], path...) + c.path = normalizePath(c.path, b2s(c.buf)) +} + +// Domain returns cookie domain. +// +// The returned value is valid until the Cookie reused or released (ReleaseCookie). +// Do not store references to the returned value. Make copies instead. +func (c *Cookie) Domain() []byte { + return c.domain +} + +// SetDomain sets cookie domain. +func (c *Cookie) SetDomain(domain string) { + c.domain = append(c.domain[:0], domain...) +} + +// SetDomainBytes sets cookie domain. +func (c *Cookie) SetDomainBytes(domain []byte) { + c.domain = append(c.domain[:0], domain...) +} + +// MaxAge returns the seconds until the cookie is meant to expire or 0 +// if no max age. +func (c *Cookie) MaxAge() int { + return c.maxAge +} + +// SetMaxAge sets cookie expiration time based on seconds. This takes precedence +// over any absolute expiry set on the cookie. +// +// Set max age to 0 to unset. +func (c *Cookie) SetMaxAge(seconds int) { + c.maxAge = seconds +} + +// Expire returns cookie expiration time. +// +// CookieExpireUnlimited is returned if cookie doesn't expire. +func (c *Cookie) Expire() time.Time { + expire := c.expire + if expire.IsZero() { + expire = cookieExpireUnlimited + } + return expire +} + +// SetExpire sets cookie expiration time. +// +// Set expiration time to CookieExpireDelete for expiring (deleting) +// the cookie on the client. +// +// By default cookie lifetime is limited by browser session. +func (c *Cookie) SetExpire(expire time.Time) { + c.expire = expire +} + +// Value returns cookie value. +// +// The returned value is valid until the Cookie reused or released (ReleaseCookie). +// Do not store references to the returned value. Make copies instead. +func (c *Cookie) Value() []byte { + return c.value +} + +// SetValue sets cookie value. +func (c *Cookie) SetValue(value string) { + c.value = append(c.value[:0], value...) +} + +// SetValueBytes sets cookie value. +func (c *Cookie) SetValueBytes(value []byte) { + c.value = append(c.value[:0], value...) +} + +// Key returns cookie name. +// +// The returned value is valid until the Cookie reused or released (ReleaseCookie). +// Do not store references to the returned value. Make copies instead. +func (c *Cookie) Key() []byte { + return c.key +} + +// SetKey sets cookie name. +func (c *Cookie) SetKey(key string) { + c.key = append(c.key[:0], key...) +} + +// SetKeyBytes sets cookie name. +func (c *Cookie) SetKeyBytes(key []byte) { + c.key = append(c.key[:0], key...) +} + +// Reset clears the cookie. +func (c *Cookie) Reset() { + c.key = c.key[:0] + c.value = c.value[:0] + c.expire = zeroTime + c.maxAge = 0 + c.domain = c.domain[:0] + c.path = c.path[:0] + c.httpOnly = false + c.secure = false + c.sameSite = CookieSameSiteDisabled +} + +// AppendBytes appends cookie representation to dst and returns +// the extended dst. +func (c *Cookie) AppendBytes(dst []byte) []byte { + if len(c.key) > 0 { + dst = append(dst, c.key...) + dst = append(dst, '=') + } + dst = append(dst, c.value...) + + if c.maxAge > 0 { + dst = append(dst, ';', ' ') + dst = append(dst, strCookieMaxAge...) + dst = append(dst, '=') + dst = appendUint(dst, c.maxAge) + } else if !c.expire.IsZero() { + dst = append(dst, ';', ' ') + dst = append(dst, strCookieExpires...) + dst = append(dst, '=') + dst = AppendHTTPDate(dst, c.expire) + } + if len(c.domain) > 0 { + dst = appendCookiePart(dst, strCookieDomain, b2s(c.domain)) + } + if len(c.path) > 0 { + dst = appendCookiePart(dst, strCookiePath, b2s(c.path)) + } + if c.httpOnly { + dst = append(dst, ';', ' ') + dst = append(dst, strCookieHTTPOnly...) + } + if c.secure { + dst = append(dst, ';', ' ') + dst = append(dst, strCookieSecure...) + } + switch c.sameSite { + case CookieSameSiteDefaultMode: + dst = append(dst, ';', ' ') + dst = append(dst, strCookieSameSite...) + case CookieSameSiteLaxMode: + dst = append(dst, ';', ' ') + dst = append(dst, strCookieSameSite...) + dst = append(dst, '=') + dst = append(dst, strCookieSameSiteLax...) + case CookieSameSiteStrictMode: + dst = append(dst, ';', ' ') + dst = append(dst, strCookieSameSite...) + dst = append(dst, '=') + dst = append(dst, strCookieSameSiteStrict...) + case CookieSameSiteNoneMode: + dst = append(dst, ';', ' ') + dst = append(dst, strCookieSameSite...) + dst = append(dst, '=') + dst = append(dst, strCookieSameSiteNone...) + } + return dst +} + +// Cookie returns cookie representation. +// +// The returned value is valid until the Cookie reused or released (ReleaseCookie). +// Do not store references to the returned value. Make copies instead. +func (c *Cookie) Cookie() []byte { + c.buf = c.AppendBytes(c.buf[:0]) + return c.buf +} + +// String returns cookie representation. +func (c *Cookie) String() string { + return string(c.Cookie()) +} + +// WriteTo writes cookie representation to w. +// +// WriteTo implements io.WriterTo interface. +func (c *Cookie) WriteTo(w io.Writer) (int64, error) { + n, err := w.Write(c.Cookie()) + return int64(n), err +} + +var errNoCookies = errors.New("no cookies found") + +// Parse parses Set-Cookie header. +func (c *Cookie) Parse(src string) error { + c.buf = append(c.buf[:0], src...) + return c.ParseBytes(c.buf) +} + +// ParseBytes parses Set-Cookie header. +func (c *Cookie) ParseBytes(src []byte) error { + c.Reset() + ntot := 0 + for { + k, v, n := parseCookie(src) + if n == 0 { + break + } else if ntot == 0 { + c.key = append(c.key, k...) + c.value = append(c.value, v...) + } + key := b2s(k) + value := b2s(v) + ntot += n + src = src[n:] + if len(key) != 0 { + // Case insensitive switch on first char + switch key[0] | 0x20 { + case 'm': + if caseInsensitiveCompare(strCookieMaxAge, key) { + maxAge, err := strconv.ParseUint(value, 10, 32) + if err != nil { + return err + } + c.maxAge = int(maxAge) + } + + case 'e': // "expires" + if caseInsensitiveCompare(strCookieExpires, key) { + + // Try the same two formats as net/http + // See: https://github.com/golang/go/blob/00379be17e63a5b75b3237819392d2dc3b313a27/src/net/http/cookie.go#L133-L135 + exptime, err := time.ParseInLocation(time.RFC1123, value, time.UTC) + if err != nil { + exptime, err = time.Parse("Mon, 02-Jan-2006 15:04:05 MST", value) + if err != nil { + return err + } + } + c.expire = exptime + } + + case 'd': // "domain" + if caseInsensitiveCompare(strCookieDomain, key) { + c.domain = append(c.domain, value...) + } + + case 'p': // "path" + if caseInsensitiveCompare(strCookiePath, key) { + c.path = append(c.path, value...) + } + + case 's': // "samesite" + if caseInsensitiveCompare(strCookieSameSite, key) { + if len(value) > 0 { + // Case insensitive switch on first char + switch value[0] | 0x20 { + case 'l': // "lax" + if caseInsensitiveCompare(strCookieSameSiteLax, value) { + c.sameSite = CookieSameSiteLaxMode + } + case 's': // "strict" + if caseInsensitiveCompare(strCookieSameSiteStrict, value) { + c.sameSite = CookieSameSiteStrictMode + } + case 'n': // "none" + if caseInsensitiveCompare(strCookieSameSiteNone, value) { + c.sameSite = CookieSameSiteNoneMode + } + } + } + } + } + } else if len(value) != 0 { + // Case insensitive switch on first char + switch value[0] | 0x20 { + case 'h': // "httponly" + if caseInsensitiveCompare(strCookieHTTPOnly, value) { + c.httpOnly = true + } + + case 's': // "secure" + if caseInsensitiveCompare(strCookieSecure, value) { + c.secure = true + } else if caseInsensitiveCompare(strCookieSameSite, value) { + c.sameSite = CookieSameSiteDefaultMode + } + } + } // else empty or no match + } + if len(c.key) == 0 && len(c.value) == 0 { + return errNoCookies + } + return nil +} + +func appendCookiePart(dst []byte, key, value string) []byte { + dst = append(dst, ';', ' ') + dst = append(dst, key...) + dst = append(dst, '=') + return append(dst, value...) +} + +func (hb *headerBuf) appendRequestCookieBytes(dst []byte) []byte { + n := len(hb.cookies) + for i := 0; i < n; i++ { + kv := hb.cookies[i] + if !kv.isValid() { + continue + } else if kv.key.len > 0 { + dst = append(dst, hb.musttoken(kv.key)...) + dst = append(dst, '=') + } + dst = append(dst, hb.musttoken(kv.value)...) + if i+1 < n { + dst = append(dst, ';', ' ') + } + } + return dst +} +func (hb *headerBuf) appendResponseCookieBytes(dst []byte) []byte { + n := len(hb.cookies) + for i := 0; i < n; i++ { + kv := hb.cookies[i] + if !kv.isValid() { + continue + } + dst = append(dst, hb.musttoken(kv.value)...) + if i+1 < n { + dst = append(dst, ';', ' ') + } + } + return dst +} + +type cookieScanner struct { + b []byte +} + +// parseCookie parses a cookie inside cookie buffer and adds it to cookie buffer.. +// +// Cookie: \r\n +func parseCookie(cookie []byte) (key, value []byte, cookieEnd int) { + if len(cookie) == 0 { + return nil, nil, 0 + } + eqIdx := bytes.IndexByte(cookie, '=') + semiIdx := bytes.IndexByte(cookie, ';') + if eqIdx > 0 && eqIdx < semiIdx { + // cookies has form key=value; + key = trimCookie(cookie[:eqIdx], false) + } else { + // cookie has no key. + eqIdx = -1 // ensure is -1. + } + if semiIdx > 0 { + // found ';' + value = trimCookie(cookie[eqIdx+1:semiIdx], true) + } else { + value = trimCookie(cookie[eqIdx+1:], true) + } + return key, value, max(len(cookie), semiIdx+1) +} + +func trimCookie(src []byte, trimQuotes bool) []byte { + for len(src) > 0 && src[0] == ' ' { + src = src[1:] // skip leading whitespace. + } + for len(src) > 0 && src[len(src)-1] == ' ' { + src = src[:len(src)-1] // skip trailing whitespace + } + if trimQuotes { + if len(src) > 1 && src[0] == '"' && src[len(src)-1] == '"' { + src = src[1 : len(src)-1] // Trim leading+trailing quotes. + } + } + return src +} + +// caseInsensitiveCompare does a case insensitive equality comparison of +// two []byte. Assumes only letters need to be matched. +func caseInsensitiveCompare(a, b string) bool { + if len(a) != len(b) { + return false + } + for i := 0; i < len(a); i++ { + if a[i]|0x20 != b[i]|0x20 { + return false + } + } + return true +} + +func normalizePath(dst []byte, src string) []byte { + dst = dst[:0] + dst = addLeadingSlash(dst, src) + dst = decodeArgAppendNoPlus(dst, src) + + // remove duplicate slashes + b := dst + bSize := len(b) + for { + n := strings.Index(b2s(b), strSlashSlash) + if n < 0 { + break + } + b = b[n:] + copy(b, b[1:]) + b = b[:len(b)-1] + bSize-- + } + dst = dst[:bSize] + + // remove /./ parts + b = dst + for { + n := strings.Index(b2s(b), strSlashDotSlash) + if n < 0 { + break + } + nn := n + len(strSlashDotSlash) - 1 + copy(b[n:], b[nn:]) + b = b[:len(b)-nn+n] + } + + // remove /foo/../ parts + for { + n := strings.Index(b2s(b), strSlashDotDotSlash) + if n < 0 { + break + } + nn := strings.LastIndexByte(b2s(b[:n]), slashChar) + if nn < 0 { + nn = 0 + } + n += len(strSlashDotDotSlash) - 1 + copy(b[nn:], b[n:]) + b = b[:len(b)-n+nn] + } + + // remove trailing /foo/.. + n := strings.LastIndex(b2s(b), strSlashDotDot) + if n >= 0 && n+len(strSlashDotDot) == len(b) { + nn := strings.LastIndexByte(b2s(b[:n]), slashChar) + if nn < 0 { + return append(dst[:0], slashChar) + } + b = b[:nn+1] + } + + if filepath.Separator == '\\' { + // remove \.\ parts + for { + n := strings.Index(b2s(b), strBackSlashDotBackSlash) + if n < 0 { + break + } + nn := n + len(strSlashDotSlash) - 1 + copy(b[n:], b[nn:]) + b = b[:len(b)-nn+n] + } + + // remove /foo/..\ parts + for { + n := strings.Index(b2s(b), strSlashDotDotBackSlash) + if n < 0 { + break + } + nn := strings.LastIndexByte(b2s(b[:n]), slashChar) + if nn < 0 { + nn = 0 + } + nn++ + n += len(strSlashDotDotBackSlash) + copy(b[nn:], b[n:]) + b = b[:len(b)-n+nn] + } + + // remove /foo\..\ parts + for { + n := strings.Index(b2s(b), strBackSlashDotDotBackSlash) + if n < 0 { + break + } + nn := strings.LastIndexByte(b2s(b[:n]), slashChar) + if nn < 0 { + nn = 0 + } + n += len(strBackSlashDotDotBackSlash) - 1 + copy(b[nn:], b[n:]) + b = b[:len(b)-n+nn] + } + + // remove trailing \foo\.. + n := strings.LastIndex(b2s(b), strBackSlashDotDot) + if n >= 0 && n+len(strSlashDotDot) == len(b) { + nn := strings.LastIndexByte(b2s(b[:n]), slashChar) + if nn < 0 { + return append(dst[:0], slashChar) + } + b = b[:nn+1] + } + } + + return b +} + +func addLeadingSlash(dst []byte, src string) []byte { + // add leading slash for unix paths + if len(src) == 0 || src[0] != slashChar { + dst = append(dst, slashChar) + } + + return dst +} + +// decodeArgAppendNoPlus is almost identical to decodeArgAppend, but it doesn't +// substitute '+' with ' '. +// +// The function is copy-pasted from decodeArgAppend due to the performance +// reasons only. +func decodeArgAppendNoPlus(dst []byte, src string) []byte { + idx := strings.IndexByte(src, '%') + if idx < 0 { + // fast path: src doesn't contain encoded chars + return append(dst, src...) + } + dst = append(dst, src[:idx]...) + + // slow path + for i := idx; i < len(src); i++ { + c := src[i] + if c == '%' { + if i+2 >= len(src) { + return append(dst, src[i:]...) + } + x2 := hex2intTable[src[i+2]] + x1 := hex2intTable[src[i+1]] + if x1 == 16 || x2 == 16 { + dst = append(dst, '%') + } else { + dst = append(dst, x1<<4|x2) + i += 2 + } + } else { + dst = append(dst, c) + } + } + return dst +} + +// AppendHTTPDate appends HTTP-compliant (RFC1123) representation of date +// to dst and returns the extended dst. +func AppendHTTPDate(dst []byte, date time.Time) []byte { + dst = date.In(time.UTC).AppendFormat(dst, time.RFC1123) + copy(dst[len(dst)-3:], strGMT) + return dst +} diff --git a/httpx/header_parse.go b/httpx/header_parse.go index 915541c..5045244 100644 --- a/httpx/header_parse.go +++ b/httpx/header_parse.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "io" + "slices" "strings" "unsafe" ) @@ -17,49 +18,37 @@ var ( errNonNumericChars = errors.New("non-numeric chars found") ) -type headerScanner struct { - b []byte - key []byte - value []byte - err error - - // hLen stores header subslice len - hLen int - - disableNormalizing bool - - // by checking whether the next line contains a colon or not to tell - // it's a header entry or a multi line value of current header entry. - // the side effect of this operation is that we know the index of the - // next colon and new line, so this can be used during next iteration, - // instead of find them again. - nextColon int - nextNewLine int - - initialized bool +func (hb *headerBuf) readFromBytes(b []byte) { + hb.buf = append(hb.buf, b...) } -type headerValueScanner struct { - b string - value string +func (hb *headerBuf) free() int { return cap(hb.buf) - len(hb.buf) } + +func (hb *headerBuf) readFrom(r io.Reader) error { + buf := hb.buf + free := hb.free() + if free == 0 { + return errSmallBuffer + } + n, err := r.Read(buf[len(buf):cap(buf)]) + hb.buf = buf[:len(buf)+n] + return err } -func (h *header) parse(buf []byte) (int, error) { - m, err := h.parseFirstLine(buf) +func (h *header) parse() (err error) { + hb := &h.hbuf + hb.off = 0 // start parsing from 0. + h.method, h.requestURI, h.proto, h.flags, err = hb.parseFirstLine(h.flags) if err != nil { - return 0, err + return err } - h.rawHeaders, _, err = readRawHeaders(h.rawHeaders[:0], b2s(buf[m:])) + var ss scannerState + err = h.parseHeaders(&ss) if err != nil { - return 0, err + return err } - var n int - n, err = h.parseHeaders(buf[m:]) - if err != nil { - return 0, err - } - return m + n, nil + return nil } func (hb *headerBuf) offBuf() []byte { @@ -87,13 +76,14 @@ func (hb *headerBuf) scanUntilByte(c byte) []byte { return buf } -func (hb *headerBuf) parseFirstLine() (method, uri, proto headerSlice, flags flags, err error) { +func (hb *headerBuf) parseFirstLine(initFlags flags) (method, uri, proto headerSlice, flags flags, err error) { var b []byte for len(b) == 0 { b = hb.scanLine() } + flags = initFlags if len(b) < 5 { - return method, uri, proto, 0, errors.New("too short first HTTP line") + return method, uri, proto, flags, errors.New("too short first HTTP line") } methodEnd := max(0, bytes.IndexByte(b, ' ')) @@ -115,6 +105,143 @@ func (hb *headerBuf) parseFirstLine() (method, uri, proto headerSlice, flags fla return method, uri, proto, flags, nil } +type scannerState struct { + err error + + // hLen stores header subslice len + hLen int + + disableNormalizing bool + + // by checking whether the next line contains a colon or not to tell + // it's a header entry or a multi line value of current header entry. + // the side effect of this operation is that we know the index of the + // next colon and new line, so this can be used during next iteration, + // instead of find them again. + nextColon int + nextNewLine int + + initialized bool +} + +func (h *header) parseHeaders(ss *scannerState) (err error) { + hb := &h.hbuf + h.contentLength = -2 + + for kv := hb.nextKV(ss); kv.isValid(); kv = hb.nextKV(ss) { + if h.flags.hasAny(disableSpecialHeader) { + h.hbuf.headers = append(h.hbuf.headers, kv) + continue + } + } + if ss.err != nil && err == nil { + err = ss.err + } + if err != nil { + h.flags |= connectionClose + return err + } + + // if h.contentLength < 0 { + // h.contentLengthBytes = hb.noKV().value + // } + if h.flags.hasAny(noHTTP11) && !h.flags.hasAny(connectionClose) { + // close connection for non-http/1.1 request unless 'Connection: keep-alive' is set. + if !h.hasHeaderValue(strConnection, strKeepAlive) { + h.flags |= connectionClose + } + } + return nil +} + +func (h *header) hasHeaderValue(key, value string) bool { + kv := h.peekHeader(key) + return kv.isValid() && b2s(h.hbuf.musttoken(kv.value)) == value +} + +func (h *header) peekHeaderBytes(key string) []byte { + kv := h.peekHeader(key) + if kv.isValid() { + return h.hbuf.musttoken(kv.value) + } + return nil +} + +// peekHeader returns header key-value for the given key. +// +// The returned value is valid until the request is released, +// either though ReleaseRequest or your request handler returning. +// Do not store references to returned value. Make copies instead. +func (h *header) peekHeader(key string) argsKV { + hb := &h.hbuf + for i := 0; i < len(h.hbuf.headers); i++ { + if b2s(hb.musttoken(h.hbuf.headers[i].key)) == key { + return h.hbuf.headers[i] + } + } + return hb.noKV() +} + +func (h *header) peekPtrHeader(key string) *argsKV { + hb := &h.hbuf + for i := 0; i < len(h.hbuf.headers); i++ { + if b2s(hb.musttoken(h.hbuf.headers[i].key)) == key { + return &h.hbuf.headers[i] + } + } + return nil +} + +func (hb *headerBuf) mustAppendSlice(value string) headerSlice { + L := len(hb.buf) + copy(hb.buf[L:L+len(value)], value) + hb.buf = hb.buf[:L+len(value)] + return hb.slice(hb.buf[L : L+len(value)]) +} + +func (h *header) reuseOrAppend(tok headerSlice, value string) headerSlice { + if tok.len > tokint(len(value)) { + copy(h.hbuf.musttoken(tok), value) + tok.len = tokint(len(value)) + return tok + } + return h.appendSlice(value) +} + +func (h *header) appendSlice(value string) headerSlice { + free := h.hbuf.free() + if len(value) > free { + if h.flags.hasAny(flagNoBufferGrow) { + h.flags |= flagOOMReached + return headerSlice{} + } + h.hbuf.buf = slices.Grow(h.hbuf.buf, len(value)) + } + return h.hbuf.mustAppendSlice(value) +} + +func (h *header) appendHeader(key, value string) { + hb := &h.hbuf + free := hb.free() + buf := h.hbuf.buf + + if len(key)+len(value) > free { + if h.flags.hasAny(flagNoBufferGrow) { + panic(errSmallBuffer) + } + slices.Grow(buf, len(key)+len(value)) + } + k := hb.mustAppendSlice(key) + v := hb.mustAppendSlice(value) + if !h.flags.hasAny(disableNormalizing) { + // TODO + } + hb.headers = append(hb.headers, argsKV{ + key: k, + value: v, + }) +} + func readRawHeaders(dst []byte, buf string) ([]byte, int, error) { n := strings.IndexByte(buf, nChar) if n < 0 { @@ -142,217 +269,125 @@ func readRawHeaders(dst []byte, buf string) ([]byte, int, error) { } } } +func (hb *headerBuf) noKV() argsKV { return argsKV{} } -func (h *header) parseHeaders(buf []byte) (int, error) { - h.contentLength = -2 - h.scanner = headerScanner{} - s := &h.scanner - s.b = buf - s.disableNormalizing = h.disableNormalizing - var err error - for s.next() { - key := b2s(s.key) - value := b2s(s.value) - if len(key) > 0 { - // Spaces between the header key and colon are not allowed. - // See RFC 7230, Section 3.2.4. - - if strings.IndexByte(key, ' ') != -1 || strings.IndexByte(key, '\t') != -1 { - err = fmt.Errorf("invalid header key %q", s.key) - continue - } - - if h.disableSpecialHeader { - h.h = appendArg(h.h, key, value, argsHasValue) - continue - } - - switch s.key[0] | 0x20 { - case 'h': - if caseInsensitiveCompare(key, strHost) { - h.host = append(h.host[:0], value...) - continue - } - case 'u': - if caseInsensitiveCompare(key, strUserAgent) { - h.userAgent = append(h.userAgent[:0], value...) - continue - } - case 'c': - if caseInsensitiveCompare(key, strContentType) { - h.contentType = append(h.contentType[:0], value...) - continue - } - if caseInsensitiveCompare(key, strContentLength) { - if h.contentLength != -1 { - var nerr error - if h.contentLength, nerr = parseContentLength(b2s(s.value)); nerr != nil { - if err == nil { - err = nerr - } - h.contentLength = -2 - } else { - h.contentLengthBytes = append(h.contentLengthBytes[:0], value...) - } - } - continue - } - if caseInsensitiveCompare(key, strConnection) { - if b2s(s.value) == strClose { - h.connectionClose = true - } else { - h.connectionClose = false - h.h = appendArg(h.h, key, value, argsHasValue) - } - continue - } - case 't': - if caseInsensitiveCompare(key, strTransferEncoding) { - if value != strIdentity { - h.contentLength = -1 - h.h = setArg(h.h, strTransferEncoding, strChunked, argsHasValue) - } - continue - } - if caseInsensitiveCompare(key, strTrailer) { - if nerr := h.SetTrailer(value); nerr != nil { - if err == nil { - err = nerr - } - } - continue - } - } - } - h.h = appendArg(h.h, key, value, argsHasValue) +func (hb *headerBuf) nextKV(ss *scannerState) argsKV { + if !ss.initialized { + ss.nextColon = -1 + ss.nextNewLine = -1 + ss.initialized = true } - if s.err != nil && err == nil { - err = s.err + buf := hb.buf[hb.off:] + bLen := len(buf) + if bLen >= 2 && buf[0] == rChar && buf[1] == nChar { + hb.off += 2 + return hb.noKV() // \r\n\r\n Ends header. } - if err != nil { - h.connectionClose = true - return 0, err - } - - if h.contentLength < 0 { - h.contentLengthBytes = h.contentLengthBytes[:0] - } - if h.noHTTP11 && !h.connectionClose { - // close connection for non-http/1.1 request unless 'Connection: keep-alive' is set. - v := peekArgStr(h.h, strConnection) - h.connectionClose = !hasHeaderValue(b2s(v), strKeepAlive) - } - return s.hLen, nil -} - -func (s *headerScanner) next() bool { - if !s.initialized { - s.nextColon = -1 - s.nextNewLine = -1 - s.initialized = true - } - bLen := len(s.b) - if bLen >= 2 && s.b[0] == rChar && s.b[1] == nChar { - s.b = s.b[2:] - s.hLen += 2 - return false - } - if bLen >= 1 && s.b[0] == nChar { - s.b = s.b[1:] - s.hLen++ - return false + if bLen >= 1 && buf[0] == nChar { + hb.off++ + return hb.noKV() // \n\n: Ends header. } var n int - if s.nextColon >= 0 { - n = s.nextColon - s.nextColon = -1 + if ss.nextColon >= 0 { + n = ss.nextColon + ss.nextColon = -1 } else { - n = bytes.IndexByte(s.b, ':') + n = bytes.IndexByte(buf, ':') // There can't be a \n inside the header name, check for this. - x := bytes.IndexByte(s.b, nChar) + x := bytes.IndexByte(buf, nChar) if x < 0 { // A header name should always at some point be followed by a \n // even if it's the one that terminates the header block. - s.err = errNeedMore - return false + ss.err = errNeedMore + return hb.noKV() } if x < n { // There was a \n before the : - s.err = errInvalidName - return false + ss.err = errInvalidName + return hb.noKV() } } if n < 0 { - s.err = errNeedMore - return false + ss.err = errNeedMore + return hb.noKV() } - s.key = s.b[:n] - normalizeHeaderKey(s.key, s.disableNormalizing) + + if bytes.IndexByte(buf[:n], ' ') >= 0 || bytes.IndexByte(buf[:n], '\t') >= 0 { + // Spaces between the header key and colon are not allowed. + // See RFC 7230, Section 3.2.4. + ss.err = errInvalidName + return hb.noKV() + } + + var resultKV argsKV + resultKV.key = hb.slice(buf[:n]) + normalizeHeaderKey(buf[:n], ss.disableNormalizing) n++ - for len(s.b) > n && s.b[n] == ' ' { + for len(buf) > n && buf[n] == ' ' { n++ // the newline index is a relative index, and lines below trimmed `s.b` by `n`, // so the relative newline index also shifted forward. it's safe to decrease // to a minus value, it means it's invalid, and will find the newline again. - s.nextNewLine-- + ss.nextNewLine-- } - s.hLen += n - s.b = s.b[n:] - if s.nextNewLine >= 0 { - n = s.nextNewLine - s.nextNewLine = -1 + ss.hLen += n + buf = buf[n:] + if ss.nextNewLine >= 0 { + n = ss.nextNewLine + ss.nextNewLine = -1 } else { - n = bytes.IndexByte(s.b, nChar) + n = bytes.IndexByte(buf, nChar) } if n < 0 { - s.err = errNeedMore - return false + ss.err = errNeedMore + return hb.noKV() } isMultiLineValue := false for { - if n+1 >= len(s.b) { + if n+1 >= len(buf) { break } - if s.b[n+1] != ' ' && s.b[n+1] != '\t' { + if buf[n+1] != ' ' && buf[n+1] != '\t' { break } - d := bytes.IndexByte(s.b[n+1:], nChar) + d := bytes.IndexByte(buf[n+1:], nChar) if d <= 0 { break - } else if d == 1 && s.b[n+1] == rChar { + } else if d == 1 && buf[n+1] == rChar { break } e := n + d + 1 - if c := bytes.IndexByte(s.b[n+1:e], ':'); c >= 0 { - s.nextColon = c - s.nextNewLine = d - c - 1 + if c := bytes.IndexByte(buf[n+1:e], ':'); c >= 0 { + ss.nextColon = c + ss.nextNewLine = d - c - 1 break } isMultiLineValue = true n = e } - if n >= len(s.b) { - s.err = errNeedMore - return false + if n >= len(buf) { + ss.err = errNeedMore + return hb.noKV() } - oldB := s.b - s.value = s.b[:n] - s.hLen += n + 1 - s.b = s.b[n+1:] + oldB := buf + value := buf[:n] + ss.hLen += n + 1 + buf = buf[n+1:] - if n > 0 && s.value[n-1] == rChar { + if n > 0 && value[n-1] == rChar { n-- } - for n > 0 && s.value[n-1] == ' ' { + for n > 0 && value[n-1] == ' ' { n-- } - s.value = s.value[:n] + value = value[:n] if isMultiLineValue { - s.value, s.b, s.hLen = normalizeHeaderValue(s.value, oldB, s.hLen) + value, buf, ss.hLen = normalizeHeaderValue(value, oldB, ss.hLen) } - return true + resultKV.value = hb.slice(value) + return resultKV } func normalizeHeaderKey(b []byte, disableNormalizing bool) { @@ -437,33 +472,6 @@ func parseContentLength(b string) (int, error) { return v, nil } -func hasHeaderValue(s, value string) bool { - var vs headerValueScanner - vs.b = s - for vs.next() { - if caseInsensitiveCompare(vs.value, value) { - return true - } - } - return false -} - -func (s *headerValueScanner) next() bool { - b := s.b - if len(b) == 0 { - return false - } - n := strings.IndexByte(b, ',') - if n < 0 { - s.value = stripSpace(b) - s.b = b[len(b):] - return true - } - s.value = stripSpace(b[:n]) - s.b = b[n+1:] - return true -} - func nextLine(b []byte) ([]byte, []byte, error) { nNext := bytes.IndexByte(b, nChar) if nNext < 0 { @@ -552,6 +560,7 @@ func (h *header) readLoop(r *bufio.Reader, waitForMore bool) error { func (h *header) tryRead(r *bufio.Reader, n int) error { h.resetSkipNormalize() b, err := r.Peek(n) + if len(b) == 0 { if err == io.EOF { return err @@ -577,35 +586,28 @@ func (h *header) tryRead(r *bufio.Reader, n int) error { return fmt.Errorf("error when reading request headers: %w", err) } b = mustPeekBuffered(r) - headersLen, errParse := h.parse(b) + errParse := h.parse() if errParse != nil { return headerError("request", err, errParse, b, false) } - mustDiscard(r, headersLen) + // mustDiscard(r, headersLen) return nil } +func (h *headerBuf) reset() { + *h = headerBuf{ + buf: h.buf[:0], + headers: h.headers[:0], + cookies: h.cookies[:0], + } +} + func (h *header) resetSkipNormalize() { - h.noHTTP11 = false - h.connectionClose = false - - h.contentLength = 0 - h.contentLengthBytes = h.contentLengthBytes[:0] - - h.method = h.method[:0] - h.proto = h.proto[:0] - h.requestURI = h.requestURI[:0] - h.host = h.host[:0] - h.contentType = h.contentType[:0] - h.userAgent = h.userAgent[:0] - h.trailer = h.trailer[:0] - h.mulHeader = h.mulHeader[:0] - - h.h = h.h[:0] - h.cookies = h.cookies[:0] - h.cookiesCollected = false - - h.rawHeaders = h.rawHeaders[:0] + h.hbuf.reset() + *h = header{ + hbuf: h.hbuf, + logger: h.logger, + } } func headerError(typ string, err, errParse error, b []byte, secureErrorLogMessage bool) error { @@ -638,6 +640,7 @@ func isOnlyCRLF(b []byte) bool { } return true } + func headerErrorMsg(typ string, err error, b []byte, secureErrorLogMessage bool) error { return fmt.Errorf("error when reading %s headers: %w. Buffer size=%d", typ, err, len(b)) } @@ -671,89 +674,19 @@ func mustDiscard(r *bufio.Reader, n int) { } } -// Peek returns header value for the given key. -// -// The returned value is valid until the request is released, -// either though ReleaseRequest or your request handler returning. -// Do not store references to returned value. Make copies instead. -func (h *header) Peek(key string) []byte { - k := getHeaderKeyBytes(&h.bufKV, key, h.disableNormalizing) - return h.peek(b2s(k)) -} - // Host returns Host header value. func (h *header) Host() []byte { - if h.disableSpecialHeader { - return peekArg(h.h, HeaderHost) - } - return h.host -} - -func getHeaderKeyBytes(kv *argsKV, key string, disableNormalizing bool) []byte { - kv.key = append(kv.key[:0], key...) - normalizeHeaderKey(kv.key, disableNormalizing) - return kv.key -} - -func peekArg(h []argsKV, k string) []byte { - for i, n := 0, len(h); i < n; i++ { - kv := &h[i] - if b2s(kv.key) == k { - return kv.value - } - } - return nil -} - -func (h *header) peek(key string) []byte { - switch key { - case HeaderHost: - return h.Host() - case HeaderContentType: - return h.ContentType() - case HeaderUserAgent: - return h.UserAgent() - case HeaderConnection: - if h.ConnectionClose() { - return []byte(strClose) - } - return peekArg(h.h, key) - case HeaderContentLength: - return h.contentLengthBytes - case HeaderCookie: - if h.cookiesCollected { - return appendRequestCookieBytes(nil, h.cookies) - } - return peekArg(h.h, key) - case HeaderTrailer: - return appendArgsKey(nil, h.trailer, strCommaSpace) - default: - return peekArg(h.h, key) - } + return h.peekHeaderBytes(HeaderHost) } // ConnectionClose returns true if 'Connection: close' header is set. func (h *header) ConnectionClose() bool { - return h.connectionClose + return h.flags.hasAny(connectionClose) } // UserAgent returns User-Agent header value. func (h *header) UserAgent() []byte { - if h.disableSpecialHeader { - return peekArg(h.h, HeaderUserAgent) - } - return h.userAgent -} - -func appendArgsKey(dst []byte, args []argsKV, sep string) []byte { - for i, n := 0, len(args); i < n; i++ { - kv := &args[i] - dst = append(dst, kv.key...) - if i+1 < n { - dst = append(dst, sep...) - } - } - return dst + return h.peekHeaderBytes(HeaderUserAgent) } // b2s converts byte slice to a string without memory allocation. diff --git a/httpx/tokenizer.go b/httpx/tokenizer.go index 667f53c..df9af6e 100644 --- a/httpx/tokenizer.go +++ b/httpx/tokenizer.go @@ -5,17 +5,19 @@ import ( "log/slog" "net/http" "strconv" - "strings" "unsafe" "github.com/soypat/lneto/internal" ) type headerBuf struct { + // buf[:len] holds entire HTTP header data, which may be normalized by [flags]. buf[off:len] holds data not yet processed during parsing. buf []byte - off int // offset into buf for parsing. + // offset into buf for parsing. + off int // args contains key-value store. - args []argsKV + headers []argsKV + cookies []argsKV } type tokint = uint16 @@ -25,38 +27,22 @@ type headerSlice struct { len tokint } -type argSlice struct { - start tokint - len tokint -} - type argsKV struct { key headerSlice value headerSlice // value start >0 means value is present. } +func (kv argsKV) isValid() bool { + return kv.key.start > 0 +} + +func (kv *argsKV) invalidate() { + *kv = argsKV{} +} + func (tb headerBuf) musttoken(slice headerSlice) []byte { return tb.buf[slice.start : slice.start+slice.len] } -func (tb headerBuf) mustargs(slice argSlice) []argsKV { - return tb.args[slice.start : slice.start+slice.len] -} - -func (tb headerBuf) visitArgs(args argSlice, f func(k, v []byte)) { - a := tb.mustargs(args) - for _, arg := range a { - k := tb.musttoken(arg.key) - v := tb.musttoken(arg.value) - f(k, v) - } -} - -func (tb headerBuf) visitArgsKey(args argSlice, f func(k []byte)) { - a := tb.mustargs(args) - for _, arg := range a { - f(tb.musttoken(arg.key)) - } -} func (tb headerBuf) slice(b []byte) headerSlice { base := uintptr(unsafe.Pointer(&tb.buf[0])) @@ -81,69 +67,37 @@ const ( connectionClose noHTTP11 cookiesCollected + flagNoBufferGrow + flagOOMReached ) -func (f flags) hasAll(checkThese flags) bool { - return f&checkThese == checkThese +func (f flags) hasAny(checkThese flags) bool { + return f&checkThese != 0 } type header struct { - buf headerBuf + hbuf headerBuf + logger *slog.Logger contentLength int - h argSlice - cookies argSlice - trailer argSlice - host headerSlice - contentLengthBytes headerSlice - contentType headerSlice - userAgent headerSlice - method headerSlice - proto headerSlice - requestURI headerSlice - rawHeaders headerSlice - mulHeader headerSlice + method headerSlice + requestURI headerSlice + proto headerSlice - flags flags - logger *slog.Logger + flags flags } func (h *header) Set(key, value string) { - h.bufKV.key = append(h.bufKV.key[:0], key...) - normalizeHeaderKey(h.bufKV.key, h.disableNormalizing) - h.SetCanonical(b2s(h.bufKV.key), value) + h.SetCanonical(key, value) //TODO: implement non-canonical. } func (h *header) Add(key, value string) { - if h.setSpecialHeader(key, value) { - return - } - k := getHeaderKeyBytes(&h.bufKV, key, h.disableNormalizing) - h.h = appendArg(h.h, b2s(k), value, argsHasValue) -} - -// ContentEncoding returns Content-Encoding header value. -func (h *header) ContentEncoding() []byte { - return peekArg(h.h, strContentEncoding) -} - -// SetContentEncoding sets Content-Encoding header value. -func (h *header) SetContentEncoding(contentEncoding string) { - h.Set(strContentEncoding, contentEncoding) -} - -// SetContentType sets Content-Type header value. -func (h *header) SetContentType(contentType string) { - h.contentType = append(h.contentType[:0], contentType...) + h.appendHeader(key, value) } // ContentType returns Content-Type header value. func (h *header) ContentType() []byte { - contentType := h.contentType - if !h.noDefaultContentType && len(h.contentType) == 0 { - contentType = append(contentType, defaultContentType...) - } - return contentType + return h.peekHeaderBytes(HeaderContentType) } // SetCanonical sets the given 'key: value' header assuming that @@ -152,109 +106,36 @@ func (h *header) ContentType() []byte { // If the header is set as a Trailer (forbidden trailers will not be set, see SetTrailer for more details), // it will be sent after the chunked request body. func (h *header) SetCanonical(key, value string) { - if h.setSpecialHeader(key, value) { - return + kv := h.peekPtrHeader(key) + if kv != nil { + kv.invalidate() } - h.setNonSpecial(key, value) -} - -// setSpecialHeader handles special headers and return true when a header is processed. -func (h *header) setSpecialHeader(key, value string) bool { - if len(key) == 0 || h.disableSpecialHeader { - return false - } - h.trace("setSpecialHeader", slog.String("key", key), slog.String("value", value)) - switch key[0] | 0x20 { - case 'c': - switch { - case caseInsensitiveCompare(strContentType, key): - h.SetContentType(value) - return true - case caseInsensitiveCompare(strContentLength, key): - if contentLength, err := parseContentLength(value); err == nil { - h.contentLength = contentLength - h.contentLengthBytes = append(h.contentLengthBytes[:0], value...) - } - return true - case caseInsensitiveCompare(strConnection, key): - if strClose == value { - h.SetConnectionClose() - } else { - h.ResetConnectionClose() - h.setNonSpecial(key, value) - } - return true - case caseInsensitiveCompare(strCookie, key): - h.collectCookies() - h.cookies = parseRequestCookies(h.cookies, value) - return true - } - case 't': // OK - if caseInsensitiveCompare(strTransferEncoding, key) { - // Transfer-Encoding is managed automatically. - return true - } else if caseInsensitiveCompare(strTrailer, key) { - _ = h.SetTrailer(value) - return true - } - case 'h': - if caseInsensitiveCompare(strHost, key) { - h.SetHost(value) - return true - } - case 'u': - if caseInsensitiveCompare(strUserAgent, key) { - h.SetUserAgent(value) - return true - } - } - - return false + h.appendHeader(key, value) } // SetHost sets Host header value. func (h *header) SetHost(host string) { - h.host = append(h.host[:0], host...) + h.Set(HeaderHost, host) } // SetUserAgent sets User-Agent header value. func (h *header) SetUserAgent(userAgent string) { - h.userAgent = append(h.userAgent[:0], userAgent...) + h.Set(HeaderUserAgent, userAgent) } // SetConnectionClose sets 'Connection: close' header. func (h *header) SetConnectionClose() { - h.connectionClose = true + h.flags |= connectionClose } // ResetConnectionClose clears 'Connection: close' header if it exists. func (h *header) ResetConnectionClose() { - if h.connectionClose { - h.connectionClose = false - h.h = delAllArgs(h.h, strConnection) + if h.flags.hasAny(connectionClose) { + h.flags &^= connectionClose + // h.h = delAllArgs(h.h, strConnection) // TODO } } -func (h *header) SetContentRange(startPos, endPos, contentLength int) { - b := h.bufKV.value[:0] - b = append(b, strBytes...) - b = append(b, ' ') - b = appendUint(b, startPos) - b = append(b, '-') - b = appendUint(b, endPos) - b = append(b, '/') - b = appendUint(b, contentLength) - h.bufKV.value = b - - h.setNonSpecial(strContentRange, b2s(h.bufKV.value)) -} - -// setNonSpecial directly put into map i.e. not a basic header. -func (h *header) setNonSpecial(key string, value string) { - h.trace("httpx:setNonSpecial", slog.String("key", key), slog.String("value", value)) - h.h = setArg(h.h, key, value, argsHasValue) -} - func appendUint(b []byte, v int) []byte { if v < 0 { panic("negative uint") @@ -271,144 +152,8 @@ func (h *header) ContentLength() int { return h.contentLength } -// SetContentLength sets Content-Length header value. -// -// Content-Length may be negative: -// -1 means Transfer-Encoding: chunked. -// -2 means Transfer-Encoding: identity. -func (h *header) SetContentLength(contentLength int) { - h.contentLength = contentLength - if contentLength >= 0 { - h.contentLengthBytes = appendUint(h.contentLengthBytes[:0], contentLength) - h.h = delAllArgs(h.h, strTransferEncoding) - } else { - h.contentLengthBytes = h.contentLengthBytes[:0] - h.h = setArg(h.h, strTransferEncoding, strChunked, argsHasValue) - } -} - var ErrBadTrailer = errors.New("contain forbidden trailer") -// SetTrailer sets Trailer header value for chunked request -// to indicate which headers will be sent after the body. -// -// Use Set to set the trailer header later. -// -// Trailers are only supported with chunked transfer. -// Trailers allow the sender to include additional headers at the end of chunked messages. -// -// The following trailers are forbidden: -// 1. necessary for message framing (e.g., Transfer-Encoding and Content-Length), -// 2. routing (e.g., Host), -// 3. request modifiers (e.g., controls and conditionals in Section 5 of [RFC7231]), -// 4. authentication (e.g., see [RFC7235] and [RFC6265]), -// 5. response control data (e.g., see Section 7.1 of [RFC7231]), -// 6. determining how to process the payload (e.g., Content-Encoding, Content-Type, Content-Range, and Trailer) -// -// Return ErrBadTrailer if contain any forbidden trailers. -func (h *header) SetTrailer(trailer string) error { - h.trailer = h.trailer[:0] - return h.AddTrailer(trailer) -} - -// AddTrailerBytes add Trailer header value for chunked response -// to indicate which headers will be sent after the body. -// -// Use Set to set the trailer header later. -// -// Trailers are only supported with chunked transfer. -// Trailers allow the sender to include additional headers at the end of chunked messages. -// -// The following trailers are forbidden: -// 1. necessary for message framing (e.g., Transfer-Encoding and Content-Length), -// 2. routing (e.g., Host), -// 3. request modifiers (e.g., controls and conditionals in Section 5 of [RFC7231]), -// 4. authentication (e.g., see [RFC7235] and [RFC6265]), -// 5. response control data (e.g., see Section 7.1 of [RFC7231]), -// 6. determining how to process the payload (e.g., Content-Encoding, Content-Type, Content-Range, and Trailer) -// -// Return ErrBadTrailer if contain any forbidden trailers. -func (h *header) AddTrailer(trailer string) error { - h.trace("httpx:AddTrailer", slog.String("trailer", trailer)) - var err error - for i := -1; i+1 < len(trailer); { - trailer = trailer[i+1:] - i = strings.IndexByte(trailer, ',') - if i < 0 { - i = len(trailer) - } - key := stripSpace(trailer[:i]) - // Forbidden by RFC 7230, section 4.1.2 - if isBadTrailer(key) { - err = ErrBadTrailer - continue - } - h.bufKV.key = append(h.bufKV.key[:0], key...) - normalizeHeaderKey(h.bufKV.key, h.disableNormalizing) - h.trailer = appendArg(h.trailer, b2s(h.bufKV.key), "", argsNoValue) - } - - return err -} - -func isBadTrailer(key string) bool { - if len(key) == 0 { - return true - } - - switch key[0] | 0x20 { - case 'a': - return caseInsensitiveCompare(key, strAuthorization) - case 'c': - if len(key) > len(HeaderContentType) && caseInsensitiveCompare(key[:8], strContentType[:8]) { - // skip compare prefix 'Content-' - return caseInsensitiveCompare(key[8:], strContentEncoding[8:]) || - caseInsensitiveCompare(key[8:], strContentLength[8:]) || - caseInsensitiveCompare(key[8:], strContentType[8:]) || - caseInsensitiveCompare(key[8:], strContentRange[8:]) - } - return caseInsensitiveCompare(key, strConnection) - case 'e': - return caseInsensitiveCompare(key, strExpect) - case 'h': - return caseInsensitiveCompare(key, strHost) - case 'k': - return caseInsensitiveCompare(key, strKeepAlive) - case 'm': - return caseInsensitiveCompare(key, strMaxForwards) - case 'p': - if len(key) > len(HeaderProxyConnection) && caseInsensitiveCompare(key[:6], strProxyConnection[:6]) { - // skip compare prefix 'Proxy-' - return caseInsensitiveCompare(key[6:], strProxyConnection[6:]) || - caseInsensitiveCompare(key[6:], strProxyAuthenticate[6:]) || - caseInsensitiveCompare(key[6:], strProxyAuthorization[6:]) - } - case 'r': - return caseInsensitiveCompare(key, strRange) - case 't': - return caseInsensitiveCompare(key, strTE) || - caseInsensitiveCompare(key, strTrailer) || - caseInsensitiveCompare(key, strTransferEncoding) - case 'w': - return caseInsensitiveCompare(key, strWWWAuthenticate) - } - return false -} - -// RawHeaders returns raw header key/value bytes. -// -// Depending on server configuration, header keys may be normalized to -// capital-case in place. -// -// This copy is set aside during parsing, so empty slice is returned for all -// cases where parsing did not happen. Similarly, request line is not stored -// during parsing and can not be returned. -// -// The slice is not safe to use after the handler returns. -func (h *header) RawHeaders() []byte { - return h.rawHeaders -} - // DisableNormalizing disables header names' normalization. // // By default all the header names are normalized by uppercasing @@ -422,90 +167,72 @@ func (h *header) RawHeaders() []byte { // // Disable header names' normalization only if know what are you doing. func (h *header) DisableNormalizing() { - h.disableNormalizing = true -} - -// DisableSpecialHeader disables special header processing. -// fasthttp will not set any special headers for you, such as Host, Content-Type, User-Agent, etc. -// You must set everything yourself. -// If RequestHeader.Read() is called, special headers will be ignored. -// This can be used to control case and order of special headers. -// This is generally not recommended. -func (h *header) DisableSpecialHeader() { - h.disableSpecialHeader = true + h.flags |= disableNormalizing } // Method returns HTTP request method. func (h *header) Method() []byte { - if len(h.method) == 0 { - h.method = append(h.method, http.MethodGet...) - } - return h.method + return h.hbuf.musttoken(h.method) } func (h *header) SetMethod(method string) { - h.method = append(h.method[:0], method...) + h.method = h.reuseOrAppend(h.method, method) } // SetRequestURI sets RequestURI for the first HTTP request line. func (h *header) SetRequestURI(requestURI string) { - h.requestURI = append(h.requestURI[:0], requestURI...) + h.requestURI = h.reuseOrAppend(h.requestURI, requestURI) } // RequestURI returns RequestURI from the first HTTP request line. func (h *header) RequestURI() []byte { - requestURI := h.requestURI - if len(requestURI) == 0 { - requestURI = append(requestURI, '/') + if h.requestURI.start == 0 { + return nil + } else if h.requestURI.len == 0 { + h.requestURI = h.appendSlice("/") } - return requestURI + return h.hbuf.musttoken(h.requestURI) } // Protocol returns HTTP protocol. func (h *header) Protocol() []byte { - if len(h.proto) == 0 { - h.proto = append(h.proto, strHTTP11...) + if h.proto.len == 0 { + h.proto = h.appendSlice(strHTTP11) } - return h.proto + return h.hbuf.musttoken(h.proto) } func (h *header) SetProtocol(protocol string) { - h.proto = append(h.proto[:0], protocol...) + h.proto = h.reuseOrAppend(h.proto, protocol) } // AppendReqRespCommon appends request/response common header representation to dst and returns the extended buffer. func (h *header) AppendReqRespCommon(dst []byte) []byte { - for i, n := 0, len(h.h); i < n; i++ { - kv := &h.h[i] - // Exclude trailer from header - exclude := false - for _, t := range h.trailer { - if b2s(kv.key) == b2s(t.key) { - exclude = true - break - } - } - if !exclude { - dst = appendHeaderLine(dst, b2s(kv.key), b2s(kv.value)) + for i, n := 0, len(h.hbuf.headers); i < n; i++ { + kv := &h.hbuf.headers[i] + if kv.isValid() { + key := h.hbuf.musttoken(kv.key) + value := h.hbuf.musttoken(kv.value) + dst = appendHeaderLine(dst, b2s(key), b2s(value)) } } - if len(h.trailer) > 0 { - aux := appendArgsKey(nil, h.trailer, strCommaSpace) - dst = appendHeaderLine(dst, strTrailer, b2s(aux)) - } + // if len(h.trailer) > 0 { + // aux := appendArgsKey(nil, h.trailer, strCommaSpace) + // dst = appendHeaderLine(dst, strTrailer, b2s(aux)) + // } // there is no need in h.collectCookies() here, since if cookies aren't collected yet, // they all are located in h.h. - n := len(h.cookies) - if n > 0 && !h.disableSpecialHeader { + n := len(h.hbuf.cookies) + if n > 0 && !h.flags.hasAny(disableSpecialHeader) { dst = append(dst, strCookie...) dst = append(dst, strColonSpace...) - dst = appendRequestCookieBytes(dst, h.cookies) + h.hbuf.appendRequestCookieBytes(dst) dst = append(dst, strCRLF...) } - if h.ConnectionClose() && !h.disableSpecialHeader { + if h.ConnectionClose() && !h.flags.hasAny(disableSpecialHeader) { dst = appendHeaderLine(dst, strConnection, strClose) } @@ -524,31 +251,37 @@ func (h *header) ignoreBody() bool { } func (h *header) collectCookies() { - if h.cookiesCollected { + if h.flags.hasAny(cookiesCollected) { return } - - for i, n := 0, len(h.h); i < n; i++ { - kv := &h.h[i] - if caseInsensitiveCompare(b2s(kv.key), strCookie) { - h.cookies = parseRequestCookies(h.cookies, b2s(kv.value)) - tmp := *kv - copy(h.h[i:], h.h[i+1:]) - n-- - i-- - h.h[n] = tmp - h.h = h.h[:n] + n := len(h.hbuf.headers) + for i := 0; i < n; i++ { + kv := h.hbuf.headers[i] + if kv.isValid() && caseInsensitiveCompare(b2s(h.hbuf.musttoken(kv.key)), HeaderCookie) { + cookie := h.hbuf.musttoken(kv.value) + for len(cookie) > 0 { + key, value, n := parseCookie(cookie) + h.hbuf.cookies = append(h.hbuf.cookies, argsKV{ + key: h.hbuf.slice(key), + value: h.hbuf.slice(value), + }) + cookie = cookie[n:] + } } } - h.cookiesCollected = true + h.flags |= cookiesCollected +} + +func (h *header) parseReqCookie(value []byte) { + } func (h *header) MethodIs(method string) bool { - return b2s(h.method) == method + return b2s(h.Method()) == method } // IsGet returns true if request method is GET. -func (h *header) IsGet() bool { return len(h.method) == 0 || h.MethodIs(http.MethodGet) } +func (h *header) IsGet() bool { return h.method.len == 0 || h.MethodIs(http.MethodGet) } // IsHead returns true if request method is HEAD. func (h *header) IsHead() bool { return h.MethodIs(http.MethodHead) } @@ -575,7 +308,7 @@ func (h *header) IsTrace() bool { return h.MethodIs(http.MethodTrace) } func (h *header) IsPatch() bool { return h.MethodIs(http.MethodPatch) } // IsHTTP11 returns true if the request is HTTP/1.1. -func (h *header) IsHTTP11() bool { return !h.noHTTP11 } +func (h *header) IsHTTP11() bool { return !h.flags.hasAny(noHTTP11) } // Embed this type into a struct, which mustn't be copied, // so `go vet` gives a warning if this struct is copied. @@ -596,27 +329,3 @@ func (h *header) debug(msg string, attrs ...slog.Attr) { func (h *header) info(msg string, attrs ...slog.Attr) { internal.LogAttrs(h.logger, slog.LevelInfo, msg, attrs...) } - -func normalizeHeaderKey(b []byte, disableNormalizing bool) { - if disableNormalizing { - return - } - - n := len(b) - if n == 0 { - return - } - - b[0] = toUpperTable[b[0]] - for i := 1; i < n; i++ { - p := &b[i] - if *p == '-' { - i++ - if i < n { - b[i] = toUpperTable[b[i]] - } - continue - } - *p = toLowerTable[*p] - } -}